diff --git a/mllib/src/main/scala/org/apache/spark/ml/Predictor.scala b/mllib/src/main/scala/org/apache/spark/ml/Predictor.scala index 17d6be0ce7cd5..38d4c43bc154b 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/Predictor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/Predictor.scala @@ -19,7 +19,7 @@ package org.apache.spark.ml import org.apache.spark.annotation.Since import org.apache.spark.internal.{LogKeys} -import org.apache.spark.ml.linalg.VectorUDT +import org.apache.spark.ml.linalg.SQLDataTypes import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util.SchemaUtils @@ -135,7 +135,7 @@ abstract class Predictor[ * * The default value is VectorUDT, but it may be overridden if FeaturesType is not Vector. */ - private[ml] def featuresDataType: DataType = new VectorUDT + private[ml] def featuresDataType: DataType = SQLDataTypes.VectorType override def transformSchema(schema: StructType): StructType = { validateAndTransformSchema(schema, fitting = true, featuresDataType) @@ -171,7 +171,7 @@ abstract class PredictionModel[FeaturesType, M <: PredictionModel[FeaturesType, * * The default value is VectorUDT, but it may be overridden if FeaturesType is not Vector. */ - protected def featuresDataType: DataType = new VectorUDT + protected def featuresDataType: DataType = SQLDataTypes.VectorType override def transformSchema(schema: StructType): StructType = { var outputSchema = validateAndTransformSchema(schema, fitting = false, featuresDataType) diff --git a/mllib/src/main/scala/org/apache/spark/ml/attribute/AttributeGroup.scala b/mllib/src/main/scala/org/apache/spark/ml/attribute/AttributeGroup.scala index f2fe125db67f9..9bb94355e65f2 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/attribute/AttributeGroup.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/attribute/AttributeGroup.scala @@ -19,7 +19,7 @@ package org.apache.spark.ml.attribute import scala.collection.mutable.ArrayBuffer -import org.apache.spark.ml.linalg.VectorUDT +import org.apache.spark.ml.linalg.SQLDataTypes import org.apache.spark.sql.types.{Metadata, MetadataBuilder, StructField} import org.apache.spark.util.ArrayImplicits._ @@ -155,7 +155,7 @@ class AttributeGroup private ( /** Converts to a StructField with some existing metadata. */ def toStructField(existingMetadata: Metadata): StructField = { - StructField(name, new VectorUDT, nullable = false, toMetadata(existingMetadata)) + StructField(name, SQLDataTypes.VectorType, nullable = false, toMetadata(existingMetadata)) } /** Converts to a StructField. */ @@ -237,7 +237,7 @@ object AttributeGroup { * Creates an attribute group from a `StructField` instance. */ def fromStructField(field: StructField): AttributeGroup = { - require(field.dataType == new VectorUDT) + require(field.dataType == SQLDataTypes.VectorType) if (field.metadata.contains(ML_ATTR)) { fromMetadata(field.metadata.getMetadata(ML_ATTR), field.name) } else { diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/Classifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/Classifier.scala index d9238479e8031..e713eec4f2d54 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/Classifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/Classifier.scala @@ -20,7 +20,7 @@ package org.apache.spark.ml.classification import org.apache.spark.annotation.Since import org.apache.spark.internal.{LogKeys} import org.apache.spark.ml.{PredictionModel, Predictor, PredictorParams} -import org.apache.spark.ml.linalg.{Vector, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector} import org.apache.spark.ml.param.ParamMap import org.apache.spark.ml.param.shared.HasRawPredictionCol import org.apache.spark.ml.util._ @@ -39,7 +39,7 @@ private[spark] trait ClassifierParams fitting: Boolean, featuresDataType: DataType): StructType = { val parentSchema = super.validateAndTransformSchema(schema, fitting, featuresDataType) - SchemaUtils.appendColumn(parentSchema, $(rawPredictionCol), new VectorUDT) + SchemaUtils.appendColumn(parentSchema, $(rawPredictionCol), SQLDataTypes.VectorType) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala index ea2c79d8a2181..925f26501edd9 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala @@ -19,7 +19,7 @@ package org.apache.spark.ml.classification import org.apache.spark.annotation.Since import org.apache.spark.internal.{LogKeys} -import org.apache.spark.ml.linalg.{DenseVector, Vector, VectorUDT} +import org.apache.spark.ml.linalg.{DenseVector, SQLDataTypes, Vector} import org.apache.spark.ml.param.ParamMap import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util.SchemaUtils @@ -37,7 +37,7 @@ private[ml] trait ProbabilisticClassifierParams fitting: Boolean, featuresDataType: DataType): StructType = { val parentSchema = super.validateAndTransformSchema(schema, fitting, featuresDataType) - SchemaUtils.appendColumn(parentSchema, $(probabilityCol), new VectorUDT) + SchemaUtils.appendColumn(parentSchema, $(probabilityCol), SQLDataTypes.VectorType) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/clustering/GaussianMixture.scala b/mllib/src/main/scala/org/apache/spark/ml/clustering/GaussianMixture.scala index 1b85f7f998751..272fce565a8f1 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/clustering/GaussianMixture.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/clustering/GaussianMixture.scala @@ -74,7 +74,7 @@ private[clustering] trait GaussianMixtureParams extends Params with HasMaxIter w protected def validateAndTransformSchema(schema: StructType): StructType = { SchemaUtils.validateVectorCompatibleColumn(schema, getFeaturesCol) val schemaWithPredictionCol = SchemaUtils.appendColumn(schema, $(predictionCol), IntegerType) - SchemaUtils.appendColumn(schemaWithPredictionCol, $(probabilityCol), new VectorUDT) + SchemaUtils.appendColumn(schemaWithPredictionCol, $(probabilityCol), SQLDataTypes.VectorType) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/clustering/LDA.scala b/mllib/src/main/scala/org/apache/spark/ml/clustering/LDA.scala index 0fc7174cadbba..1f072bd852d2a 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/clustering/LDA.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/clustering/LDA.scala @@ -355,7 +355,7 @@ private[clustering] trait LDAParams extends Params with HasFeaturesCol with HasM } } SchemaUtils.validateVectorCompatibleColumn(schema, getFeaturesCol) - SchemaUtils.appendColumn(schema, $(topicDistributionCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(topicDistributionCol), SQLDataTypes.VectorType) } private[clustering] def getOldOptimizer: OldLDAOptimizer = diff --git a/mllib/src/main/scala/org/apache/spark/ml/evaluation/BinaryClassificationEvaluator.scala b/mllib/src/main/scala/org/apache/spark/ml/evaluation/BinaryClassificationEvaluator.scala index 1a97eb2910056..8183cd34365d3 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/evaluation/BinaryClassificationEvaluator.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/evaluation/BinaryClassificationEvaluator.scala @@ -18,7 +18,7 @@ package org.apache.spark.ml.evaluation import org.apache.spark.annotation.Since -import org.apache.spark.ml.linalg.{Vector, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector} import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util._ @@ -115,7 +115,8 @@ class BinaryClassificationEvaluator @Since("1.4.0") (@Since("1.4.0") override va @Since("3.1.0") def getMetrics(dataset: Dataset[_]): BinaryClassificationMetrics = { val schema = dataset.schema - SchemaUtils.checkColumnTypes(schema, $(rawPredictionCol), Seq(DoubleType, new VectorUDT)) + SchemaUtils.checkColumnTypes(schema, $(rawPredictionCol), + Seq(DoubleType, SQLDataTypes.VectorType)) SchemaUtils.checkNumericType(schema, $(labelCol)) if (isDefined(weightCol)) { SchemaUtils.checkNumericType(schema, $(weightCol)) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/Binarizer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/Binarizer.scala index 52ed90415f1cd..50952010581fe 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/Binarizer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/Binarizer.scala @@ -218,7 +218,7 @@ final class Binarizer @Since("1.4.0") (@Since("1.4.0") override val uid: String) SchemaUtils.getSchemaField(schema, inputColName) ).size if (size < 0) { - StructField(outputColName, new VectorUDT) + StructField(outputColName, SQLDataTypes.VectorType) } else { new AttributeGroup(outputColName, numAttributes = size).toStructField() } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/BucketedRandomProjectionLSH.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/BucketedRandomProjectionLSH.scala index 3d1765417775e..ec7726a5d8104 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/BucketedRandomProjectionLSH.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/BucketedRandomProjectionLSH.scala @@ -205,7 +205,7 @@ class BucketedRandomProjectionLSH(override val uid: String) @Since("2.1.0") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) validateAndTransformSchema(schema) } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/CountVectorizer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/CountVectorizer.scala index a85e9236517f7..033c53d5afd76 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/CountVectorizer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/CountVectorizer.scala @@ -24,7 +24,7 @@ import org.apache.spark.annotation.Since import org.apache.spark.broadcast.Broadcast import org.apache.spark.ml.{Estimator, Model} import org.apache.spark.ml.attribute.{Attribute, AttributeGroup, NumericAttribute} -import org.apache.spark.ml.linalg.{Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vectors} import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared.{HasInputCol, HasOutputCol} import org.apache.spark.ml.util._ @@ -99,7 +99,7 @@ private[feature] trait CountVectorizerParams extends Params with HasInputCol wit protected def validateAndTransformSchema(schema: StructType): StructType = { val typeCandidates = List(new ArrayType(StringType, true), new ArrayType(StringType, false)) SchemaUtils.checkColumnTypes(schema, $(inputCol), typeCandidates) - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } /** diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/DCT.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/DCT.scala index 9a8bfb195666b..fe240557c0fae 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/DCT.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/DCT.scala @@ -22,7 +22,7 @@ import org.jtransforms.dct._ import org.apache.spark.annotation.Since import org.apache.spark.ml.UnaryTransformer import org.apache.spark.ml.attribute.AttributeGroup -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors, VectorUDT} import org.apache.spark.ml.param.BooleanParam import org.apache.spark.ml.util._ import org.apache.spark.sql.types._ @@ -71,10 +71,11 @@ class DCT @Since("1.5.0") (@Since("1.5.0") override val uid: String) override protected def validateInputType(inputType: DataType): Unit = { require(inputType.isInstanceOf[VectorUDT], - s"Input type must be ${(new VectorUDT).catalogString} but got ${inputType.catalogString}.") + s"Input type must be ${SQLDataTypes.VectorType.catalogString} " + + s"but got ${inputType.catalogString}.") } - override protected def outputDataType: DataType = new VectorUDT + override protected def outputDataType: DataType = SQLDataTypes.VectorType override def transformSchema(schema: StructType): StructType = { var outputSchema = super.transformSchema(schema) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/ElementwiseProduct.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/ElementwiseProduct.scala index 6dac09d13b99b..4e38a91cdb660 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/ElementwiseProduct.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/ElementwiseProduct.scala @@ -79,10 +79,11 @@ class ElementwiseProduct @Since("1.4.0") (@Since("1.4.0") override val uid: Stri override protected def validateInputType(inputType: DataType): Unit = { require(inputType.isInstanceOf[VectorUDT], - s"Input type must be ${(new VectorUDT).catalogString} but got ${inputType.catalogString}.") + s"Input type must be ${SQLDataTypes.VectorType.catalogString} " + + s"but got ${inputType.catalogString}.") } - override protected def outputDataType: DataType = new VectorUDT() + override protected def outputDataType: DataType = SQLDataTypes.VectorType override def transformSchema(schema: StructType): StructType = { var outputSchema = super.transformSchema(schema) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/IDF.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/IDF.scala index 3ef916bc200c0..3d677cddf8feb 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/IDF.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/IDF.scala @@ -60,8 +60,8 @@ private[feature] trait IDFBase extends Params with HasInputCol with HasOutputCol * Validate and transform the input schema. */ protected def validateAndTransformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/Interaction.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/Interaction.scala index 3311231e6d830..7fa1f8e6a1de4 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/Interaction.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/Interaction.scala @@ -23,7 +23,7 @@ import org.apache.spark.SparkException import org.apache.spark.annotation.Since import org.apache.spark.ml.Transformer import org.apache.spark.ml.attribute._ -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors, VectorUDT} import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util._ @@ -64,7 +64,7 @@ class Interaction @Since("1.6.0") (@Since("1.6.0") override val uid: String) ext require(get(outputCol).isDefined, "Output col must be defined first.") require($(inputCols).length > 0, "Input cols must have non-zero length.") require($(inputCols).distinct.length == $(inputCols).length, "Input cols must be distinct.") - StructType(schema.fields :+ StructField($(outputCol), new VectorUDT, false)) + StructType(schema.fields :+ StructField($(outputCol), SQLDataTypes.VectorType, false)) } @Since("2.0.0") diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/LSH.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/LSH.scala index 9c3b39b12bdc6..4aa58efcbc9c4 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/LSH.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/LSH.scala @@ -18,7 +18,7 @@ package org.apache.spark.ml.feature import org.apache.spark.ml.{Estimator, Model} -import org.apache.spark.ml.linalg.{Vector, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector} import org.apache.spark.ml.param.{IntParam, ParamValidators} import org.apache.spark.ml.param.shared.{HasInputCol, HasOutputCol} import org.apache.spark.ml.util._ @@ -53,7 +53,8 @@ private[ml] trait LSHParams extends HasInputCol with HasOutputCol { * @return A derived schema with [[outputCol]] added. */ protected[this] final def validateAndTransformSchema(schema: StructType): StructType = { - SchemaUtils.appendColumn(schema, $(outputCol), DataTypes.createArrayType(new VectorUDT)) + SchemaUtils.appendColumn(schema, $(outputCol), + DataTypes.createArrayType(SQLDataTypes.VectorType)) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/MaxAbsScaler.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/MaxAbsScaler.scala index a962ce3784c9a..05e38f7e70489 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/MaxAbsScaler.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/MaxAbsScaler.scala @@ -23,7 +23,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.annotation.Since import org.apache.spark.ml.{Estimator, Model} -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.param.{ParamMap, Params} import org.apache.spark.ml.param.shared.{HasInputCol, HasOutputCol} import org.apache.spark.ml.stat.Summarizer @@ -39,10 +39,10 @@ private[feature] trait MaxAbsScalerParams extends Params with HasInputCol with H /** Validates and transforms the input schema. */ protected def validateAndTransformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) require(!schema.fieldNames.contains($(outputCol)), s"Output column ${$(outputCol)} already exists.") - val outputFields = schema.fields :+ StructField($(outputCol), new VectorUDT, false) + val outputFields = schema.fields :+ StructField($(outputCol), SQLDataTypes.VectorType, false) StructType(outputFields) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/MinHashLSH.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/MinHashLSH.scala index 4b0e5ca4fb31a..46ef7dca720dd 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/MinHashLSH.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/MinHashLSH.scala @@ -24,7 +24,7 @@ import scala.util.Random import org.apache.hadoop.fs.Path import org.apache.spark.annotation.Since -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.param.ParamMap import org.apache.spark.ml.param.shared.HasSeed import org.apache.spark.ml.util._ @@ -201,7 +201,7 @@ class MinHashLSH(override val uid: String) extends LSH[MinHashLSHModel] with Has @Since("2.1.0") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) validateAndTransformSchema(schema) } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/MinMaxScaler.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/MinMaxScaler.scala index 9bf13c48b22aa..339710a6f280d 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/MinMaxScaler.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/MinMaxScaler.scala @@ -23,7 +23,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.annotation.Since import org.apache.spark.ml.{Estimator, Model} -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.param.{DoubleParam, ParamMap, Params} import org.apache.spark.ml.param.shared.{HasInputCol, HasOutputCol} import org.apache.spark.ml.stat.Summarizer @@ -64,10 +64,10 @@ private[feature] trait MinMaxScalerParams extends Params with HasInputCol with H /** Validates and transforms the input schema. */ protected def validateAndTransformSchema(schema: StructType): StructType = { require($(min) < $(max), s"The specified min(${$(min)}) is larger or equal to max(${$(max)})") - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) require(!schema.fieldNames.contains($(outputCol)), s"Output column ${$(outputCol)} already exists.") - val outputFields = schema.fields :+ StructField($(outputCol), new VectorUDT, false) + val outputFields = schema.fields :+ StructField($(outputCol), SQLDataTypes.VectorType, false) StructType(outputFields) } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/Normalizer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/Normalizer.scala index c7b7164e42f36..83b6092edc2d0 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/Normalizer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/Normalizer.scala @@ -20,7 +20,7 @@ package org.apache.spark.ml.feature import org.apache.spark.annotation.Since import org.apache.spark.ml.UnaryTransformer import org.apache.spark.ml.attribute.AttributeGroup -import org.apache.spark.ml.linalg.{Vector, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, VectorUDT} import org.apache.spark.ml.param.{DoubleParam, ParamValidators} import org.apache.spark.ml.util._ import org.apache.spark.mllib.feature @@ -62,10 +62,11 @@ class Normalizer @Since("1.4.0") (@Since("1.4.0") override val uid: String) override protected def validateInputType(inputType: DataType): Unit = { require(inputType.isInstanceOf[VectorUDT], - s"Input type must be ${(new VectorUDT).catalogString} but got ${inputType.catalogString}.") + s"Input type must be ${SQLDataTypes.VectorType.catalogString} " + + s"but got ${inputType.catalogString}.") } - override protected def outputDataType: DataType = new VectorUDT() + override protected def outputDataType: DataType = SQLDataTypes.VectorType @Since("1.4.0") override def transformSchema(schema: StructType): StructType = { diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/PCA.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/PCA.scala index a731b6750573c..f6067f03e76fa 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/PCA.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/PCA.scala @@ -51,7 +51,7 @@ private[feature] trait PCAParams extends Params with HasInputCol with HasOutputC /** Validates and transforms the input schema. */ protected def validateAndTransformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) require(!schema.fieldNames.contains($(outputCol)), s"Output column ${$(outputCol)} already exists.") SchemaUtils.updateAttributeGroupSize(schema, $(outputCol), $(k)) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/PolynomialExpansion.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/PolynomialExpansion.scala index 592ca001a2467..aa61b2335487e 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/PolynomialExpansion.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/PolynomialExpansion.scala @@ -70,10 +70,11 @@ class PolynomialExpansion @Since("1.4.0") (@Since("1.4.0") override val uid: Str override protected def validateInputType(inputType: DataType): Unit = { require(inputType.isInstanceOf[VectorUDT], - s"Input type must be ${(new VectorUDT).catalogString} but got ${inputType.catalogString}.") + s"Input type must be ${SQLDataTypes.VectorType.catalogString} " + + s"but got ${inputType.catalogString}.") } - override protected def outputDataType: DataType = new VectorUDT() + override protected def outputDataType: DataType = SQLDataTypes.VectorType @Since("1.4.1") override def copy(extra: ParamMap): PolynomialExpansion = defaultCopy(extra) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/RFormula.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/RFormula.scala index 844cedb4a302d..3b760d1783a14 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/RFormula.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/RFormula.scala @@ -27,7 +27,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.annotation.Since import org.apache.spark.ml.{Estimator, Model, Pipeline, PipelineModel, PipelineStage, Transformer} import org.apache.spark.ml.attribute.AttributeGroup -import org.apache.spark.ml.linalg.{Vector, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, VectorUDT} import org.apache.spark.ml.param.{BooleanParam, Param, ParamMap, ParamValidators} import org.apache.spark.ml.param.shared.{HasFeaturesCol, HasHandleInvalid, HasLabelCol} import org.apache.spark.ml.util._ @@ -314,9 +314,9 @@ class RFormula @Since("1.5.0") (@Since("1.5.0") override val uid: String) require(!hasLabelCol(schema) || !$(forceIndexLabel), "If label column already exists, forceIndexLabel can not be set with true.") if (hasLabelCol(schema)) { - StructType(schema.fields :+ StructField($(featuresCol), new VectorUDT, true)) + StructType(schema.fields :+ StructField($(featuresCol), SQLDataTypes.VectorType, true)) } else { - StructType(schema.fields :+ StructField($(featuresCol), new VectorUDT, true) :+ + StructType(schema.fields :+ StructField($(featuresCol), SQLDataTypes.VectorType, true) :+ StructField($(labelCol), DoubleType, true)) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/RobustScaler.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/RobustScaler.scala index 0b520542566c8..7a44076b6037c 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/RobustScaler.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/RobustScaler.scala @@ -92,10 +92,10 @@ private[feature] trait RobustScalerParams extends Params with HasInputCol with H protected def validateAndTransformSchema(schema: StructType): StructType = { require($(lower) < $(upper), s"The specified lower quantile(${$(lower)}) is " + s"larger or equal to upper quantile(${$(upper)})") - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) require(!schema.fieldNames.contains($(outputCol)), s"Output column ${$(outputCol)} already exists.") - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/Selector.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/Selector.scala index 1914a98014daa..a654c74c3582e 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/Selector.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/Selector.scala @@ -257,9 +257,9 @@ private[ml] abstract class Selector[T <: SelectorModel[T]] @Since("3.1.0") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) SchemaUtils.checkNumericType(schema, $(labelCol)) - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } @Since("3.1.0") @@ -301,7 +301,7 @@ private[ml] abstract class SelectorModel[T <: SelectorModel[T]] ( @Since("3.1.0") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) val newField = SelectorModel.prepOutputField(schema, selectedFeatures, $(outputCol), $(featuresCol), isNumericAttribute) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/StandardScaler.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/StandardScaler.scala index fd61753c25dda..d21384787fe52 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/StandardScaler.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StandardScaler.scala @@ -62,10 +62,10 @@ private[feature] trait StandardScalerParams extends Params with HasInputCol with /** Validates and transforms the input schema. */ protected def validateAndTransformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) require(!schema.fieldNames.contains($(outputCol)), s"Output column ${$(outputCol)} already exists.") - val outputFields = schema.fields :+ StructField($(outputCol), new VectorUDT, false) + val outputFields = schema.fields :+ StructField($(outputCol), SQLDataTypes.VectorType, false) StructType(outputFields) } diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/UnivariateFeatureSelector.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/UnivariateFeatureSelector.scala index 6eb5696520cda..088b1086014bd 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/UnivariateFeatureSelector.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/UnivariateFeatureSelector.scala @@ -26,7 +26,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.annotation.Since import org.apache.spark.ml.{Estimator, Model} import org.apache.spark.ml.attribute.{Attribute, AttributeGroup, NominalAttribute, NumericAttribute} -import org.apache.spark.ml.linalg.{DenseVector, SparseVector, Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{DenseVector, SparseVector, SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared.{HasFeaturesCol, HasLabelCol, HasOutputCol} import org.apache.spark.ml.stat.{ANOVATest, ChiSquareTest, FValueTest} @@ -266,9 +266,9 @@ final class UnivariateFeatureSelector @Since("3.1.1")(@Since("3.1.1") override v } } require(isSet(featureType) && isSet(labelType), "featureType and labelType need to be set") - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) SchemaUtils.checkNumericType(schema, $(labelCol)) - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } @Since("3.1.1") @@ -320,7 +320,7 @@ class UnivariateFeatureSelectorModel private[ml]( @Since("3.1.1") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) val newField = UnivariateFeatureSelectorModel .prepOutputField(schema, selectedFeatures, $(outputCol), $(featuresCol), isNumericAttribute) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/VarianceThresholdSelector.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/VarianceThresholdSelector.scala index cdbdf122dc05e..23b9a4bc67b75 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/VarianceThresholdSelector.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/VarianceThresholdSelector.scala @@ -104,8 +104,8 @@ with DefaultParamsWritable { @Since("3.1.0") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } @Since("3.1.0") @@ -159,7 +159,7 @@ class VarianceThresholdSelectorModel private[ml]( @Since("3.1.0") override def transformSchema(schema: StructType): StructType = { - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) val newField = SelectorModel.prepOutputField(schema, selectedFeatures, $(outputCol), $(featuresCol), true) SchemaUtils.appendColumn(schema, newField) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/VectorAssembler.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/VectorAssembler.scala index 831a8a33afecb..5d4156b69065e 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/VectorAssembler.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/VectorAssembler.scala @@ -25,7 +25,7 @@ import org.apache.spark.SparkException import org.apache.spark.annotation.Since import org.apache.spark.ml.Transformer import org.apache.spark.ml.attribute.{Attribute, AttributeGroup, NumericAttribute, UnresolvedAttribute} -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors, VectorUDT} import org.apache.spark.ml.param.{Param, ParamMap, ParamValidators} import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util._ @@ -173,7 +173,7 @@ class VectorAssembler @Since("1.4.0") (@Since("1.4.0") override val uid: String) if (schema.fieldNames.contains(outputColName)) { throw new IllegalArgumentException(s"Output column $outputColName already exists.") } - StructType(schema.fields :+ new StructField(outputColName, new VectorUDT, true)) + StructType(schema.fields :+ new StructField(outputColName, SQLDataTypes.VectorType, true)) } @Since("1.4.1") diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/VectorIndexer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/VectorIndexer.scala index 57d28bba2eeb5..d1810d8092ed4 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/VectorIndexer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/VectorIndexer.scala @@ -29,7 +29,7 @@ import org.apache.spark.SparkException import org.apache.spark.annotation.Since import org.apache.spark.ml.{Estimator, Model} import org.apache.spark.ml.attribute._ -import org.apache.spark.ml.linalg.{DenseVector, SparseVector, Vector, VectorUDT} +import org.apache.spark.ml.linalg.{DenseVector, SparseVector, SQLDataTypes, Vector} import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util._ @@ -160,11 +160,10 @@ class VectorIndexer @Since("1.4.0") ( override def transformSchema(schema: StructType): StructType = { // We do not transfer feature metadata since we do not know what types of features we will // produce in transform(). - val dataType = new VectorUDT require(isDefined(inputCol), s"VectorIndexer requires input column parameter: $inputCol") require(isDefined(outputCol), s"VectorIndexer requires output column parameter: $outputCol") - SchemaUtils.checkColumnType(schema, $(inputCol), dataType) - SchemaUtils.appendColumn(schema, $(outputCol), dataType) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } @Since("1.4.1") @@ -461,12 +460,11 @@ class VectorIndexerModel private[ml] ( @Since("1.4.0") override def transformSchema(schema: StructType): StructType = { - val dataType = new VectorUDT require(isDefined(inputCol), s"VectorIndexerModel requires input column parameter: $inputCol") require(isDefined(outputCol), s"VectorIndexerModel requires output column parameter: $outputCol") - SchemaUtils.checkColumnType(schema, $(inputCol), dataType) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) // If the input metadata specifies numFeatures, compare with expected numFeatures. val origAttrGroup = AttributeGroup.fromStructField( diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/VectorSlicer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/VectorSlicer.scala index 58a44a41f0e84..c5a49cf24374f 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/VectorSlicer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/VectorSlicer.scala @@ -150,7 +150,7 @@ final class VectorSlicer @Since("1.5.0") (@Since("1.5.0") override val uid: Stri override def transformSchema(schema: StructType): StructType = { require($(indices).length > 0 || $(names).length > 0, s"VectorSlicer requires that at least one feature be selected.") - SchemaUtils.checkColumnType(schema, $(inputCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(inputCol), SQLDataTypes.VectorType) if (schema.fieldNames.contains($(outputCol))) { throw new IllegalArgumentException(s"Output column ${$(outputCol)} already exists.") diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/Word2Vec.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/Word2Vec.scala index 2baf6260479f5..ae1039fc81c09 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/Word2Vec.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/Word2Vec.scala @@ -24,7 +24,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.annotation.Since import org.apache.spark.internal.config.Kryo.KRYO_SERIALIZER_MAX_BUFFER_SIZE import org.apache.spark.ml.{Estimator, Model} -import org.apache.spark.ml.linalg.{BLAS, Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{BLAS, SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util._ @@ -113,7 +113,7 @@ private[feature] trait Word2VecBase extends Params protected def validateAndTransformSchema(schema: StructType): StructType = { val typeCandidates = List(new ArrayType(StringType, true), new ArrayType(StringType, false)) SchemaUtils.checkColumnTypes(schema, $(inputCol), typeCandidates) - SchemaUtils.appendColumn(schema, $(outputCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(outputCol), SQLDataTypes.VectorType) } } diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/AFTSurvivalRegression.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/AFTSurvivalRegression.scala index d96500ea84ab9..d573d7d524a8d 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/AFTSurvivalRegression.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/AFTSurvivalRegression.scala @@ -111,14 +111,14 @@ private[regression] trait AFTSurvivalRegressionParams extends PredictorParams protected def validateAndTransformSchema( schema: StructType, fitting: Boolean): StructType = { - SchemaUtils.checkColumnType(schema, $(featuresCol), new VectorUDT) + SchemaUtils.checkColumnType(schema, $(featuresCol), SQLDataTypes.VectorType) if (fitting) { SchemaUtils.checkNumericType(schema, $(censorCol)) SchemaUtils.checkNumericType(schema, $(labelCol)) } val schemaWithQuantilesCol = if (hasQuantilesCol) { - SchemaUtils.appendColumn(schema, $(quantilesCol), new VectorUDT) + SchemaUtils.appendColumn(schema, $(quantilesCol), SQLDataTypes.VectorType) } else schema SchemaUtils.appendColumn(schemaWithQuantilesCol, $(predictionCol), DoubleType) diff --git a/mllib/src/main/scala/org/apache/spark/ml/source/libsvm/LibSVMRelation.scala b/mllib/src/main/scala/org/apache/spark/ml/source/libsvm/LibSVMRelation.scala index 7b0de1176fa80..bc11fb9a64b19 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/source/libsvm/LibSVMRelation.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/source/libsvm/LibSVMRelation.scala @@ -27,7 +27,7 @@ import org.apache.spark.TaskContext import org.apache.spark.internal.Logging import org.apache.spark.ml.attribute.AttributeGroup import org.apache.spark.ml.feature.LabeledPoint -import org.apache.spark.ml.linalg.{Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vectors, VectorUDT} import org.apache.spark.mllib.util.MLUtils import org.apache.spark.sql.{Row, SparkSession} import org.apache.spark.sql.catalyst.InternalRow @@ -82,7 +82,7 @@ private[libsvm] case class LibSVMFileFormat() if ( dataSchema.size != 2 || !DataTypeUtils.sameType(dataSchema(0).dataType, DataTypes.DoubleType) || - !DataTypeUtils.sameType(dataSchema(1).dataType, new VectorUDT()) || + !DataTypeUtils.sameType(dataSchema(1).dataType, SQLDataTypes.VectorType) || !(forWriting || dataSchema(1).metadata.getLong(LibSVMOptions.NUM_FEATURES).toInt > 0) ) { throw new IOException(s"Illegal schema for libsvm data, schema=$dataSchema") diff --git a/mllib/src/main/scala/org/apache/spark/ml/stat/ANOVATest.scala b/mllib/src/main/scala/org/apache/spark/ml/stat/ANOVATest.scala index 2a3470e38f6ef..c1ac69b6fcaed 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/stat/ANOVATest.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/stat/ANOVATest.scala @@ -97,7 +97,7 @@ private[ml] object ANOVATest { val spark = dataset.sparkSession import spark.implicits._ - SchemaUtils.checkColumnType(dataset.schema, featuresCol, new VectorUDT) + SchemaUtils.checkColumnType(dataset.schema, featuresCol, SQLDataTypes.VectorType) SchemaUtils.checkNumericType(dataset.schema, labelCol) val points = dataset.select(col(labelCol).cast("double"), col(featuresCol)) diff --git a/mllib/src/main/scala/org/apache/spark/ml/stat/ChiSquareTest.scala b/mllib/src/main/scala/org/apache/spark/ml/stat/ChiSquareTest.scala index cdbfb6090acf5..14b9e6a2be334 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/stat/ChiSquareTest.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/stat/ChiSquareTest.scala @@ -18,7 +18,7 @@ package org.apache.spark.ml.stat import org.apache.spark.annotation.Since -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.util.SchemaUtils import org.apache.spark.mllib.linalg.{Vectors => OldVectors} import org.apache.spark.mllib.stat.test.{ChiSqTest => OldChiSqTest} @@ -71,7 +71,7 @@ object ChiSquareTest { featuresCol: String, labelCol: String, flatten: Boolean): DataFrame = { - SchemaUtils.checkColumnType(dataset.schema, featuresCol, new VectorUDT) + SchemaUtils.checkColumnType(dataset.schema, featuresCol, SQLDataTypes.VectorType) SchemaUtils.checkNumericType(dataset.schema, labelCol) val spark = dataset.sparkSession diff --git a/mllib/src/main/scala/org/apache/spark/ml/stat/FValueTest.scala b/mllib/src/main/scala/org/apache/spark/ml/stat/FValueTest.scala index 56b7c058a5379..16f25621654b1 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/stat/FValueTest.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/stat/FValueTest.scala @@ -20,7 +20,7 @@ package org.apache.spark.ml.stat import org.apache.commons.math3.distribution.FDistribution import org.apache.spark.annotation.Since -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors} import org.apache.spark.ml.util.SchemaUtils import org.apache.spark.rdd.RDD import org.apache.spark.sql.{DataFrame, Dataset, Row} @@ -99,7 +99,7 @@ private[ml] object FValueTest { dataset: Dataset[_], featuresCol: String, labelCol: String): RDD[(Int, Double, Long, Double)] = { - SchemaUtils.checkColumnType(dataset.schema, featuresCol, new VectorUDT) + SchemaUtils.checkColumnType(dataset.schema, featuresCol, SQLDataTypes.VectorType) SchemaUtils.checkNumericType(dataset.schema, labelCol) val spark = dataset.sparkSession diff --git a/mllib/src/main/scala/org/apache/spark/ml/stat/Summarizer.scala b/mllib/src/main/scala/org/apache/spark/ml/stat/Summarizer.scala index 5a136b7578f36..ddda2418225ae 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/stat/Summarizer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/stat/Summarizer.scala @@ -22,7 +22,7 @@ import java.io._ import org.apache.spark.annotation.Since import org.apache.spark.internal.Logging import org.apache.spark.ml.feature.Instance -import org.apache.spark.ml.linalg.{Vector, Vectors, VectorUDT} +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector, Vectors, VectorUDT} import org.apache.spark.rdd.RDD import org.apache.spark.sql.Column import org.apache.spark.sql.catalyst.InternalRow @@ -302,7 +302,7 @@ private[spark] object SummaryBuilderImpl extends Logging { } } - private val vectorUDT = new VectorUDT + private val vectorUDT = SQLDataTypes.VectorType.asInstanceOf[VectorUDT] /** * All the metrics that can be currently computed by Spark for vectors. diff --git a/mllib/src/main/scala/org/apache/spark/ml/tree/treeParams.scala b/mllib/src/main/scala/org/apache/spark/ml/tree/treeParams.scala index 2244d49b2a35f..08dd0039a9913 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/tree/treeParams.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/tree/treeParams.scala @@ -24,7 +24,7 @@ import scala.util.Try import org.apache.spark.annotation.Since import org.apache.spark.ml.PredictorParams import org.apache.spark.ml.classification.ProbabilisticClassifierParams -import org.apache.spark.ml.linalg.VectorUDT +import org.apache.spark.ml.linalg.SQLDataTypes import org.apache.spark.ml.param._ import org.apache.spark.ml.param.shared._ import org.apache.spark.ml.util.SchemaUtils @@ -420,7 +420,7 @@ private[ml] trait TreeEnsembleClassifierParams featuresDataType: DataType): StructType = { var outputSchema = super.validateAndTransformSchema(schema, fitting, featuresDataType) if ($(leafCol).nonEmpty) { - outputSchema = SchemaUtils.appendColumn(outputSchema, $(leafCol), new VectorUDT) + outputSchema = SchemaUtils.appendColumn(outputSchema, $(leafCol), SQLDataTypes.VectorType) } outputSchema } @@ -438,7 +438,7 @@ private[ml] trait TreeEnsembleRegressorParams featuresDataType: DataType): StructType = { var outputSchema = super.validateAndTransformSchema(schema, fitting, featuresDataType) if ($(leafCol).nonEmpty) { - outputSchema = SchemaUtils.appendColumn(outputSchema, $(leafCol), new VectorUDT) + outputSchema = SchemaUtils.appendColumn(outputSchema, $(leafCol), SQLDataTypes.VectorType) } outputSchema } diff --git a/mllib/src/main/scala/org/apache/spark/ml/util/MetadataUtils.scala b/mllib/src/main/scala/org/apache/spark/ml/util/MetadataUtils.scala index 631261af249f2..2e0dbe75dee7a 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/util/MetadataUtils.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/util/MetadataUtils.scala @@ -20,7 +20,7 @@ package org.apache.spark.ml.util import scala.collection.immutable.HashMap import org.apache.spark.ml.attribute._ -import org.apache.spark.ml.linalg.VectorUDT +import org.apache.spark.ml.linalg.{SQLDataTypes, VectorUDT} import org.apache.spark.sql.types.StructField @@ -46,7 +46,7 @@ private[spark] object MetadataUtils { * Returns None if the number of features is not specified. */ def getNumFeatures(vectorSchema: StructField): Option[Int] = { - if (vectorSchema.dataType == new VectorUDT) { + if (vectorSchema.dataType == SQLDataTypes.VectorType) { val group = AttributeGroup.fromStructField(vectorSchema) val size = group.size if (size >= 0) { diff --git a/mllib/src/main/scala/org/apache/spark/ml/util/SchemaUtils.scala b/mllib/src/main/scala/org/apache/spark/ml/util/SchemaUtils.scala index 5386641838726..3bb249f201b15 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/util/SchemaUtils.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/util/SchemaUtils.scala @@ -19,7 +19,7 @@ package org.apache.spark.ml.util import org.apache.spark.SparkIllegalArgumentException import org.apache.spark.ml.attribute._ -import org.apache.spark.ml.linalg.VectorUDT +import org.apache.spark.ml.linalg.SQLDataTypes import org.apache.spark.sql.catalyst.util.{AttributeNameParser, QuotingUtils} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ @@ -201,7 +201,8 @@ private[spark] object SchemaUtils { * @param colName column name */ def validateVectorCompatibleColumn(schema: StructType, colName: String): Unit = { - val typeCandidates = List( new VectorUDT, + val typeCandidates = List( + SQLDataTypes.VectorType, new ArrayType(DoubleType, false), new ArrayType(FloatType, false)) checkColumnTypes(schema, colName, typeCandidates)