diff --git a/core/src/main/python/synapse/ml/stages/EnsembleByKey.py b/core/src/main/python/synapse/ml/stages/EnsembleByKey.py new file mode 100644 index 00000000000..4de593ac0b2 --- /dev/null +++ b/core/src/main/python/synapse/ml/stages/EnsembleByKey.py @@ -0,0 +1,13 @@ +# Copyright (C) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. See LICENSE in project root for information. + +from pyspark.ml.common import inherit_doc +from synapse.ml.stages._EnsembleByKey import _EnsembleByKey + + +@inherit_doc +class EnsembleByKey(_EnsembleByKey): + def getColNames(self): + if self.isSet(self.colNames): + return self.getOrDefault(self.colNames) + return [f"{self.getStrategy()}({name})" for name in self.getCols()] diff --git a/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.scala b/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.scala index 72484947c4f..47ea2bb7724 100644 --- a/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.scala +++ b/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.scala @@ -11,17 +11,83 @@ import org.apache.spark.ml.linalg.SQLDataTypes._ import org.apache.spark.ml.param._ import org.apache.spark.ml.stat.Summarizer import org.apache.spark.ml.util.{DefaultParamsReadable, DefaultParamsWritable, Identifiable} +import org.apache.spark.sql.catalyst.analysis.UnresolvedAttribute +import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, ExprId, RowOrdering} import org.apache.spark.sql.functions._ import org.apache.spark.sql.types._ -import org.apache.spark.sql.{DataFrame, Dataset} +import org.apache.spark.sql.{Column, DataFrame, Dataset, SparkSession} import scala.collection.JavaConverters._ - -object EnsembleByKey extends DefaultParamsReadable[EnsembleByKey] +import scala.util.Try + +object EnsembleByKey extends DefaultParamsReadable[EnsembleByKey] { + + // Spark's union analysis re-aliases duplicated child outputs and tags them with this key so that + // AttributeSeq.resolve can prune them before reporting an ambiguous reference. + private val DuplicateMetadataKey = "__is_duplicate" + + private case class PathStep(name: String, mapKeyType: Option[DataType]) + + private case class ResolvedField( + reference: String, + qualifier: Array[String], + path: Array[PathStep], + ordinals: Array[Int], + field: StructField) + + private case class ResolvedColumns( + inputFields: Array[ResolvedField], + outputNames: Array[String], + keyFields: Array[ResolvedField], + aggregateFields: Array[StructField], + caseSensitive: Boolean) + + private case class ResolvedStep( + fieldName: String, + dataType: DataType, + nullable: Boolean, + metadata: Metadata, + ordinal: Int, + mapKeyType: Option[DataType]) + + private case class FieldRole( + consumesInputColumn: Boolean, + declaredOutput: StructField => Option[StructField]) + + private case class QualifiedMatch( + qualifier: Array[String], + requestedPath: Array[String], + ordinal: Int, + exprId: ExprId) + + private def columnNamesMatch(left: String, right: String, caseSensitive: Boolean): Boolean = + if (caseSensitive) left == right else left.equalsIgnoreCase(right) + + private def resolveFieldAtLevel( + schema: StructType, + fieldName: String, + reference: String, + caseSensitive: Boolean + ): (StructField, Int) = { + schema.fields.zipWithIndex.filter { case (field, _) => + columnNamesMatch(field.name, fieldName, caseSensitive) + } match { + case Array(result) => result + case Array() => throw new IllegalArgumentException( + s"$reference does not exist. Available: ${schema.fieldNames.mkString(", ")}") + case matches => throw new IllegalArgumentException( + s"$reference is ambiguous. Matches: ${matches.map(_._1.name).mkString(", ")}") + } + } +} class EnsembleByKey(val uid: String) extends Transformer with Wrappable with DefaultParamsWritable with SynapseMLLogging { + + import EnsembleByKey._ + logClass(FeatureNames.Core) + override protected lazy val pyInternalWrapper = true def this() = this(Identifiable.randomUID("EnsembleByKey")) @@ -47,7 +113,7 @@ class EnsembleByKey(val uid: String) extends Transformer val colNames = new StringArrayParam(this, "colNames", "Names of the result of each col") - def getColNames: Array[String] = $(colNames) + def getColNames: Array[String] = get(colNames).getOrElse(getCols.map(name => s"$getStrategy($name)")) def setColNames(arr: Array[String]): this.type = set(colNames, arr) @@ -83,73 +149,645 @@ class EnsembleByKey(val uid: String) extends Transformer setDefault(collapseGroup -> true) - override def transform(dataset: Dataset[_]): DataFrame = { - logTransform[DataFrame]({ + private val aggregateType: DataType => Option[DataType] = { + case _: DoubleType => Some(DoubleType) + case _: FloatType => Some(DoubleType) + case fdt if fdt == VectorType => Some(VectorType) + case _ => None + } - if (get(colNames).isEmpty) { - setDefault(colNames -> getCols.map(name => s"$getStrategy($name)")) - } + private val aggregateField = (outputName: String, dataType: DataType) => + StructField(outputName, dataType, nullable = dataType != VectorType) - transformSchema(dataset.schema) + private val keyRole = FieldRole( + consumesInputColumn = true, + field => Some(field.copy(name = ""))) + + private val aggregateRole = FieldRole( + consumesInputColumn = false, + field => aggregateType(field.dataType).map(aggregateField("", _))) + + private val topLevelMatches = (schema: StructType, fieldName: String, caseSensitive: Boolean) => + schema.fields.zipWithIndex.collect { + case (field, ordinal) if columnNamesMatch(field.name, fieldName, caseSensitive) => ordinal + } - val strategyToFloatFunction = Map( - "mean" -> { (x: String, y: String) => mean(x).alias(y) } - ) + private val analyzedAttributes = (dataset: Option[Dataset[_]]) => + dataset.toSeq.flatMap(_.queryExecution.analyzed.output) - val strategyToVectorFunction = Map( - "mean" -> { (x: String, y: String) => - Summarizer.mean(col(x)).alias(y) + private def pruneDuplicates[A](candidates: Seq[A])(metadataOf: A => Metadata): Seq[A] = { + if (candidates.length <= 1) { + candidates + } else { + val pruned = candidates.filterNot(metadataOf(_).contains(DuplicateMetadataKey)) + if (pruned.isEmpty) candidates else pruned + } + } + + private val withoutDuplicateMarker = (metadata: Metadata) => + if (metadata.contains(DuplicateMetadataKey)) { + new MetadataBuilder().withMetadata(metadata).remove(DuplicateMetadataKey).build() + } else { + metadata + } + + private val declaredField = (field: StructField, name: String) => + field.copy(name = name, metadata = withoutDuplicateMarker(field.metadata)) + + private val shareOneExpression = (attributes: Seq[Attribute], ordinals: Array[Int]) => + ordinals.length > 1 && ordinals.forall(_ < attributes.length) && + ordinals.map(attributes(_).exprId).distinct.length == 1 + + private def qualifiedPathMatches( + attributes: Seq[Attribute], + parsedPath: Array[String], + caseSensitive: Boolean + ): Seq[QualifiedMatch] = { + attributes.zipWithIndex + .flatMap { case (attribute, ordinal) => + (1 until parsedPath.length).collect { + case index + if columnNamesMatch(attribute.name, parsedPath(index), caseSensitive) && + qualifiersMatch(attribute.qualifier, parsedPath.take(index), caseSensitive) => + QualifiedMatch(parsedPath.take(index), parsedPath.drop(index), ordinal, attribute.exprId) } - ) - - val newCols = getCols.zip(getColNames).map { case (inColName, outColName) => - dataset.schema(inColName).dataType match { - case _: DoubleType => - strategyToFloatFunction(getStrategy)(inColName, outColName) - case _: FloatType => - strategyToFloatFunction(getStrategy)(inColName, outColName) - case v if v == VectorType => - strategyToVectorFunction(getStrategy)(inColName, outColName) - case t => - throw new IllegalArgumentException(s"Cannot operate on type $t with strategy $getStrategy") + } + } + + private def qualifiedMatch( + parsedPath: Array[String], + reference: String, + caseSensitive: Boolean, + dataset: Option[Dataset[_]] + ): Option[QualifiedMatch] = { + val attributes = analyzedAttributes(dataset) + val allMatches = qualifiedPathMatches(attributes, parsedPath, caseSensitive) + if (allMatches.isEmpty) { + None + } else { + // Spark selects the qualifier/name candidate set first and only then prunes duplicate-marked + // candidates, so pruning must never change which qualifier length wins. + val longestQualifier = allMatches.map(_.qualifier.length).max + val selected = allMatches.filter(_.qualifier.length == longestQualifier) + val matches = pruneDuplicates(selected)(candidate => attributes(candidate.ordinal).metadata) + require( + matches.map(_.exprId).distinct.length == 1, + s"$reference is ambiguous because it matches multiple dataset attributes") + Some(matches.head) + } + } + + private val schemaSplit = (schema: StructType, parsedPath: Array[String], caseSensitive: Boolean) => + parsedPath.indices.filter(index => + schema.fields.exists(field => + columnNamesMatch(field.name, parsedPath(index), caseSensitive))) + + private val qualifiersMatch = (actual: Seq[String], configured: Array[String], caseSensitive: Boolean) => + actual.length >= configured.length && + actual.takeRight(configured.length).zip(configured) + .forall { case (left, right) => columnNamesMatch(left, right, caseSensitive) } + + private def bindQualifier( + dataset: Dataset[_], + resolved: ResolvedField, + caseSensitive: Boolean + ): ResolvedField = { + if (resolved.qualifier.isEmpty) { + resolved + } else { + val candidates = dataset.queryExecution.analyzed.output.zipWithIndex.filter { case (attribute, _) => + columnNamesMatch(attribute.name, resolved.path.head.name, caseSensitive) && + qualifiersMatch(attribute.qualifier, resolved.qualifier, caseSensitive) + } + val matches = pruneDuplicates(candidates)(_._1.metadata) + matches match { + case Seq() => + throw new IllegalArgumentException(s"${resolved.reference} does not match a dataset qualifier") + case _ if matches.map(_._1.exprId).distinct.length == 1 => + resolved.copy(ordinals = resolved.ordinals.updated(0, matches.head._2)) + case _ => + throw new IllegalArgumentException(s"${resolved.reference} is ambiguous") + } + } + } + + // Spark's GetMapValue casts the requested literal to the map key type and additionally requires + // that key type to be orderable (TypeUtils.checkForOrderingExpr -> RowOrdering.isOrderable). + // RowOrdering.isOrderable(DataType) is identical in Spark 3.5 and Spark 4.1, so it is safe here. + private val mapKeyIsExtractable = (keyType: DataType) => + Cast.canCast(StringType, keyType) && RowOrdering.isOrderable(keyType) + + private val unsupportedMapKeyMessage = (reference: String, keyType: DataType) => + s"$reference cannot be extracted because map key type $keyType " + ( + if (Cast.canCast(StringType, keyType)) { + "is not orderable, so Spark cannot look up a map value by key. " + + "Use a map column whose key type is orderable, such as string." + } else "does not accept string keys") + + private def resolveStep( + currentType: DataType, + fieldName: String, + currentNullable: Boolean, + reference: String, + caseSensitive: Boolean + ): ResolvedStep = { + currentType match { + case currentSchema: StructType => + val (field, fieldOrdinal) = resolveFieldAtLevel( + currentSchema, + fieldName, + reference, + caseSensitive) + ResolvedStep( + field.name, + field.dataType, + currentNullable || field.nullable, + field.metadata, + fieldOrdinal, + None) + case ArrayType(elementSchema: StructType, containsNull) => + val (field, fieldOrdinal) = + resolveFieldAtLevel(elementSchema, fieldName, reference, caseSensitive) + ResolvedStep( + field.name, + ArrayType(field.dataType, containsNull || field.nullable), + currentNullable, + Metadata.empty, + fieldOrdinal, + None) + case MapType(keyType, valueType, _) if mapKeyIsExtractable(keyType) => + ResolvedStep(fieldName, valueType, nullable = true, Metadata.empty, -1, Some(keyType)) + case MapType(keyType, _, _) => + throw new IllegalArgumentException(unsupportedMapKeyMessage(reference, keyType)) + case _ => + throw new IllegalArgumentException( + s"$reference is not supported by Spark nested field extraction") + } + } + + private def resolvePath( + currentType: DataType, + remainingPath: List[String], + currentNullable: Boolean, + ordinals: List[Int], + reference: String, + caseSensitive: Boolean + ): (StructField, List[Int], List[PathStep]) = { + val step = resolveStep( + currentType, + remainingPath.head, + currentNullable, + reference, + caseSensitive) + val pathStep = PathStep(step.fieldName, step.mapKeyType) + + remainingPath.tail match { + case Nil => + (StructField(step.fieldName, step.dataType, step.nullable, step.metadata), + ordinals :+ step.ordinal, + List(pathStep)) + case nestedPath => + val (field, fieldOrdinals, fieldSteps) = resolvePath( + step.dataType, + nestedPath, + step.nullable, + ordinals :+ step.ordinal, + reference, + caseSensitive) + (field, fieldOrdinals, pathStep +: fieldSteps) + } + } + + private def candidateOutput( + schema: StructType, + requestedPath: Array[String], + ordinal: Int, + reference: String, + caseSensitive: Boolean, + role: FieldRole + ): Option[Option[StructField]] = { + Try(resolveAtOrdinal(schema, Array.empty[String], requestedPath, ordinal, reference, caseSensitive)) + .toOption + .map(resolved => role.declaredOutput(resolved.field)) + } + + private val candidateOutputsAgree = ( + schema: StructType, + matches: Array[Int], + requestedPath: Array[String], + reference: String, + caseSensitive: Boolean, + role: FieldRole) => + matches + .map(ordinal => candidateOutput(schema, requestedPath, ordinal, reference, caseSensitive, role)) + .distinct + .length <= 1 + + private def requireStableQualifiedField( + schema: StructType, + matches: Array[Int], + requestedPath: Array[String], + reference: String, + caseSensitive: Boolean, + role: FieldRole + ): Unit = { + require( + candidateOutputsAgree(schema, matches, requestedPath, reference, caseSensitive, role), + s"$reference matches columns with incompatible declared outputs") + require( + matches.length <= 1 || requestedPath.length > 1 || + getCollapseGroup || !role.consumesInputColumn, + s"$reference cannot be resolved from schema because multiple columns are named " + + s"${requestedPath.head} when collapseGroup is false") + } + + private def resolveFromSchema( + schema: StructType, + qualifier: Array[String], + requestedPath: Array[String], + reference: String, + caseSensitive: Boolean, + role: FieldRole, + dataset: Option[Dataset[_]] + ): ResolvedField = { + val candidates = topLevelMatches(schema, requestedPath.head, caseSensitive) + if (qualifier.isEmpty) { + resolveUnqualifiedFromSchema(schema, candidates, requestedPath, reference, caseSensitive, role, dataset) + } else if (candidates.isEmpty) { + resolveNestedPath(schema, qualifier, requestedPath, reference, caseSensitive) + } else { + // A schema carries no qualifier metadata, so every ordinal the dataset-aware path could select + // must derive the same output instead of pruning duplicate-marked fields out of the candidates. + requireStableQualifiedField(schema, candidates, requestedPath, reference, caseSensitive, role) + resolveAtOrdinal(schema, qualifier, requestedPath, candidates.head, reference, caseSensitive) + } + } + + private def resolveUnqualifiedFromSchema( + schema: StructType, + candidates: Array[Int], + requestedPath: Array[String], + reference: String, + caseSensitive: Boolean, + role: FieldRole, + dataset: Option[Dataset[_]] + ): ResolvedField = { + val matches = pruneDuplicates(candidates.toSeq)(schema(_).metadata).toArray + val resolvableDuplicates = candidates.length > 1 && + (matches.length == 1 || + ((dataset.isEmpty || shareOneExpression(analyzedAttributes(dataset), matches)) && + Try(requireStableQualifiedField( + schema, matches, requestedPath, reference, caseSensitive, role)).isSuccess)) + if (!resolvableDuplicates) { + resolveNestedPath(schema, Array.empty[String], requestedPath, reference, caseSensitive) + } else { + requireStableQualifiedField(schema, matches, requestedPath, reference, caseSensitive, role) + resolveAtOrdinal(schema, Array.empty[String], requestedPath, matches.head, reference, caseSensitive) + } + } + + private def resolveNestedPath( + schema: StructType, + qualifier: Array[String], + requestedPath: Array[String], + reference: String, + caseSensitive: Boolean + ): ResolvedField = { + val (field, ordinals, steps) = + resolvePath(schema, requestedPath.toList, false, Nil, reference, caseSensitive) + ResolvedField(reference, qualifier, steps.toArray, ordinals.toArray, declaredField(field, requestedPath.last)) + } + + private def resolveFromOrdinal( + schema: StructType, + qualifier: Array[String], + requestedPath: Array[String], + ordinal: Int, + reference: String, + caseSensitive: Boolean, + role: FieldRole + ): ResolvedField = { + val matches = topLevelMatches(schema, requestedPath.head, caseSensitive) + requireStableQualifiedField(schema, matches, requestedPath, reference, caseSensitive, role) + resolveAtOrdinal(schema, qualifier, requestedPath, ordinal, reference, caseSensitive) + } + + private def resolveAtOrdinal( + schema: StructType, + qualifier: Array[String], + requestedPath: Array[String], + ordinal: Int, + reference: String, + caseSensitive: Boolean + ): ResolvedField = { + val topField = schema(ordinal) + val topStep = PathStep(topField.name, None) + val (field, ordinals, steps) = requestedPath.tail.toList match { + case Nil => (topField, List(ordinal), List(topStep)) + case nestedPath => + val (nestedField, nestedOrdinals, nestedSteps) = resolvePath( + topField.dataType, + nestedPath, + topField.nullable, + List(ordinal), + reference, + caseSensitive) + (nestedField, nestedOrdinals, topStep +: nestedSteps) + } + ResolvedField(reference, qualifier, steps.toArray, ordinals.toArray, declaredField(field, requestedPath.last)) + } + + private def outputContribution( + role: FieldRole, + resolved: ResolvedField + ): (Option[StructField], Option[Int]) = { + val consumesOrdinal = + role.consumesInputColumn && !getCollapseGroup && resolved.path.length == 1 + val consumedOrdinal = if (consumesOrdinal) Some(resolved.ordinals.head) else None + (role.declaredOutput(resolved.field), consumedOrdinal) + } + + private def schemaInterpretations( + schema: StructType, + parsedPath: Array[String], + reference: String, + caseSensitive: Boolean, + role: FieldRole, + dataset: Option[Dataset[_]] + ): Seq[ResolvedField] = { + schemaSplit(schema, parsedPath, caseSensitive).flatMap(index => + Try(resolveFromSchema( + schema, + parsedPath.take(index), + parsedPath.drop(index), + reference, + caseSensitive, + role, + dataset)).toOption) + } + + private def resolveField( + schema: StructType, + reference: String, + caseSensitive: Boolean, + dataset: Option[Dataset[_]], + role: FieldRole + ): ResolvedField = { + val parsedPath = UnresolvedAttribute.parseAttributeName(reference).toArray + val interpretations = schemaInterpretations(schema, parsedPath, reference, caseSensitive, role, dataset) + require( + interpretations.map(outputContribution(role, _)).distinct.length <= 1, + s"$reference is ambiguous between a nested field and a dataset qualifier") + + qualifiedMatch(parsedPath, reference, caseSensitive, dataset) match { + case Some(matched) => + resolveFromOrdinal( + schema, + matched.qualifier, + matched.requestedPath, + matched.ordinal, + reference, + caseSensitive, + role) + case None => + interpretations.headOption.getOrElse { + val pathStart = schemaSplit(schema, parsedPath, caseSensitive).headOption.getOrElse(0) + resolveFromSchema( + schema, + parsedPath.take(pathStart), + parsedPath.drop(pathStart), + reference, + caseSensitive, + role, + dataset) } + } + } + + private def validateNonCollapsedKeys( + schema: StructType, + keyFields: Array[ResolvedField], + outputNames: Array[String], + caseSensitive: Boolean + ): Unit = { + val keyOutputCollisions = outputNames.filter(outputName => + keyFields.exists(resolved => + columnNamesMatch(resolved.field.name, outputName, caseSensitive))).distinct + require( + keyOutputCollisions.isEmpty, + s"Output columns ${keyOutputCollisions.mkString(", ")} cannot overwrite grouping keys " + + s"${keyFields.map(_.field.name).mkString(", ")} when collapseGroup is false") + + val nestedKeyCollisions = keyFields.filter(_.path.length > 1).filter(resolved => + schema.fields.exists(field => + columnNamesMatch(field.name, resolved.field.name, caseSensitive))) + require( + nestedKeyCollisions.isEmpty, + s"Nested grouping keys ${nestedKeyCollisions.map(_.reference).mkString(", ")} " + + "cannot overwrite top-level columns when collapseGroup is false") + + val duplicateNestedKeyNames = keyFields.indices.flatMap { leftIndex => + ((leftIndex + 1) until keyFields.length).collect { + case rightIndex + if columnNamesMatch( + keyFields(leftIndex).field.name, + keyFields(rightIndex).field.name, + caseSensitive) => + keyFields(leftIndex).field.name + } + }.distinct + require( + duplicateNestedKeyNames.isEmpty, + s"Grouping keys must resolve to distinct output columns when collapseGroup is false: " + + duplicateNestedKeyNames.mkString(", ")) + } + + private def getSchemaFields( + schema: StructType, + dataset: Option[Dataset[_]] = None + ): ResolvedColumns = { + val inputNames = get(cols).getOrElse( + throw new IllegalArgumentException("cols must be set and non-empty")) + val keyNames = get(keys).getOrElse( + throw new IllegalArgumentException("keys must be set and non-empty")) + require(inputNames.nonEmpty, "cols must be set and non-empty") + require(keyNames.nonEmpty, "keys must be set and non-empty") + val outputNames = get(colNames).getOrElse( + inputNames.map(name => s"$getStrategy($name)")) + require( + inputNames.length == outputNames.length, + s"cols (${inputNames.length}) and colNames (${outputNames.length}) must have the same length") + + val caseSensitive = dataset.map(_.sparkSession).orElse(SparkSession.getActiveSession) + .exists(_.conf.get("spark.sql.caseSensitive", "false").trim.toBoolean) + val inputFields = inputNames.map(resolveField(schema, _, caseSensitive, dataset, aggregateRole)) + val keyFields = keyNames.map(resolveField(schema, _, caseSensitive, dataset, keyRole)) + keyFields.foreach { key => + require(RowOrdering.isOrderable(key.field.dataType), + s"${key.reference} resolves to ${key.field.dataType}, which Spark cannot use as a grouping key") + } + if (!getCollapseGroup) { + validateNonCollapsedKeys(schema, keyFields, outputNames, caseSensitive) + } + + val aggregateFields = inputFields.zip(outputNames).map { case (resolvedInput, outputName) => + aggregateType(resolvedInput.field.dataType) + .map(aggregateField(outputName, _)) + .getOrElse(throw new IllegalArgumentException( + s"Cannot operate on type ${resolvedInput.field.dataType} with strategy $getStrategy")) + } + + ResolvedColumns(inputFields, outputNames, keyFields, aggregateFields, caseSensitive) + } + + private def bindQualifiers( + dataset: Dataset[_], + resolvedColumns: ResolvedColumns + ): ResolvedColumns = { + resolvedColumns.copy( + inputFields = resolvedColumns.inputFields.map(bindQualifier( + dataset, + _, + resolvedColumns.caseSensitive)), + keyFields = resolvedColumns.keyFields.map(bindQualifier( + dataset, + _, + resolvedColumns.caseSensitive))) + } + + private val quoteIdentifier = (name: String) => s"`${name.replace("`", "``")}`" + + private val inputName = (index: Int) => s"__ensemble_by_key_input_$index" + + private val keyName = (index: Int) => s"__ensemble_by_key_key_$index" + + private val aggregateName = (index: Int) => s"__ensemble_by_key_aggregate_$index" + + private val normalize = (dataset: Dataset[_]) => + dataset.toDF(dataset.schema.indices.map(inputName): _*) + + private def resolvedColumn(resolved: ResolvedField): Column = { + val root = col(quoteIdentifier(inputName(resolved.ordinals.head))) + resolved.path.tail.foldLeft(root) { (column, step) => + step.mapKeyType match { + case Some(keyType) => column(lit(step.name).cast(keyType)) + case None => column.getField(step.name) } + } + } - val aggregated = dataset.toDF() - .groupBy(getKeys.head, getKeys.tail: _*) - .agg(newCols.head, newCols.tail: _*) + // The identity cast prevents grouping analysis from propagating source metadata to the key. + private val keyColumn = (resolved: ResolvedField, index: Int) => + resolvedColumn(resolved).cast(resolved.field.dataType).as(keyName(index), resolved.field.metadata) + + private def aggregateColumn( + resolvedInput: ResolvedField, + outputName: String + ): Column = { + val inputColumn = resolvedColumn(resolvedInput) + aggregateType(resolvedInput.field.dataType) match { + case Some(fdt) if fdt == VectorType => Summarizer.mean(inputColumn).alias(outputName) + case Some(_) => mean(inputColumn).alias(outputName) + case None => throw new IllegalArgumentException( + s"Cannot operate on type ${resolvedInput.field.dataType} with strategy $getStrategy") + } + } + + private def aggregate( + dataset: Dataset[_], + normalized: DataFrame, + resolvedColumns: ResolvedColumns + ): DataFrame = { + val keyColumns = resolvedColumns.keyFields.zipWithIndex.map { case (r, i) => keyColumn(r, i) } + val newColumns = resolvedColumns.inputFields.zipWithIndex.map { case (resolvedInput, index) => + aggregateColumn(resolvedInput, aggregateName(index)) + } + val retainGroupColumns = dataset.sparkSession.conf + .get("spark.sql.retainGroupColumns", "true").trim.toBoolean + val aggregateColumns = if (retainGroupColumns) newColumns else keyColumns ++ newColumns + + normalized + .groupBy(keyColumns: _*) + .agg(aggregateColumns.head, aggregateColumns.tail: _*) + } + + private def outputKeyColumns(resolvedColumns: ResolvedColumns): Array[Column] = { + resolvedColumns.keyFields.zipWithIndex.map { case (resolved, index) => + col(quoteIdentifier(keyName(index))).as(resolved.field.name, resolved.field.metadata) + } + } + + private def outputAggregateColumns(resolvedColumns: ResolvedColumns): Array[Column] = { + resolvedColumns.outputNames.indices.map(index => + col(quoteIdentifier(aggregateName(index))).as(resolvedColumns.outputNames(index))).toArray + } + + private def passthroughColumns( + schema: StructType, + resolvedColumns: ResolvedColumns + ): Array[Column] = { + val topLevelKeyOrdinals = resolvedColumns.keyFields.filter(_.path.length == 1) + .map(_.ordinals.head).toSet + schema.fields.zipWithIndex.collect { + case (field, index) + if !topLevelKeyOrdinals(index) && + !resolvedColumns.outputNames.exists(outputName => + columnNamesMatch(field.name, outputName, resolvedColumns.caseSensitive)) => + col(quoteIdentifier(inputName(index))).as(field.name, field.metadata) + } + } + + private def mergeWithGroups( + normalized: DataFrame, + aggregated: DataFrame, + resolvedColumns: ResolvedColumns, + inputSchema: StructType + ): DataFrame = { + val leftKeys = resolvedColumns.keyFields.zipWithIndex.map { case (r, i) => keyColumn(r, i) } + val left = normalized.select((col("*") +: leftKeys.toSeq): _*) + val conditions = resolvedColumns.keyFields.indices.map(i => left(keyName(i)) <=> aggregated(keyName(i))) + val joined = left.join(aggregated, conditions.reduce(_ && _)).select( + (left.columns.map(left(_)) ++ resolvedColumns.outputNames.indices.map(i => + aggregated(aggregateName(i)))): _*) + val outputColumns = + outputKeyColumns(resolvedColumns) ++ + passthroughColumns(inputSchema, resolvedColumns) ++ + outputAggregateColumns(resolvedColumns) + joined.select(outputColumns: _*) + } + + override def transform(dataset: Dataset[_]): DataFrame = { + logTransform[DataFrame]({ + val resolvedColumns = bindQualifiers(dataset, getSchemaFields(dataset.schema, Some(dataset))) + val normalized = normalize(dataset) + val aggregated = aggregate(dataset, normalized, resolvedColumns) if (getCollapseGroup) { - aggregated + aggregated.select((outputKeyColumns(resolvedColumns) ++ + outputAggregateColumns(resolvedColumns)): _*) } else { - val needToDrop = getColNames.toSet & dataset.columns.toSet - dataset.drop(needToDrop.toList: _*).toDF().join(aggregated, getKeys) + mergeWithGroups(normalized, aggregated, resolvedColumns, dataset.schema) } }, dataset.columns.length) - } def transformSchema(schema: StructType): StructType = { - val colSet = getCols.toSet - val colToNewName = getCols.zip(getColNames).toMap - - val newFields = schema.fields.flatMap { f => - if (!colSet(f.name)) None - else { - val newField = StructField(colToNewName(f.name), f.dataType) - f.dataType match { - case _: DoubleType => Some(newField) - case _: FloatType => Some(newField) - case fdt if fdt == VectorType => Some(newField) - case t => throw new IllegalArgumentException(s"Cannot operate on type $t with strategy $getStrategy") - } + val resolvedColumns = getSchemaFields(schema) + val fields = if (getCollapseGroup) { + resolvedColumns.keyFields.map(_.field) ++ resolvedColumns.aggregateFields + } else { + val topLevelKeyOrdinals = resolvedColumns.keyFields.filter(_.path.length == 1) + .map(_.ordinals.head).toSet + val inputFields = schema.fields.zipWithIndex.collect { + case (field, index) + if !topLevelKeyOrdinals(index) && + !resolvedColumns.outputNames.exists(outputName => + columnNamesMatch(field.name, outputName, resolvedColumns.caseSensitive)) => + field } + resolvedColumns.keyFields.map(_.field) ++ inputFields ++ resolvedColumns.aggregateFields } - val keyFields = schema.fields.filter(f => colSet(f.name)) - val fields = - (if (getCollapseGroup) schema.fields else keyFields).++(newFields) - new StructType(fields) } diff --git a/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.txt b/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.txt index 52d490f4f86..1d5f14fa329 100644 --- a/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.txt +++ b/core/src/main/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKey.txt @@ -5,3 +5,21 @@ the first row of the column. To avoid materialization you can provide the vector through the ``setVectorDims`` function, which takes a mapping from columns (String) to dimension (Int). You can also choose to squash or keep the original dataset with the ``collapseGroup`` parameter. + +Column references support Spark field syntax, including dataset qualifiers, nested +struct paths, array-of-struct extraction, map extraction (the referenced segment is +cast from a string literal to the map key type, following Spark cast rules, and the +map key type must also be orderable because Spark looks map values up by key), and +backtick-quoted literal field names. Duplicate columns that Spark treats as one +expression resolve like a single column, union columns that Spark marks as duplicates +are pruned the same way Spark prunes them (only within the candidate set the requested +qualifier and name already selected, so a duplicate-marked ``u.group`` still wins over an +untagged ``v.group``), while references that match several distinct attributes are +rejected as ambiguous. Because a ``StructType`` does not retain dataset aliases, +``transformSchema`` cannot reject a qualifier that matches no dataset; ``transform`` +detects and reports that invalid qualifier when the analyzed dataset is available. +Schema-only case resolution similarly uses the active Spark session, while runtime +resolution uses the dataset session. If no matching active session exists and those +sessions use different ``spark.sql.caseSensitive`` values, pipeline schema validation +can differ from runtime resolution; keep the dataset session active while constructing +or validating a pipeline. diff --git a/core/src/test/python/synapsemltest/stages/__init__.py b/core/src/test/python/synapsemltest/stages/__init__.py new file mode 100644 index 00000000000..f780f4fea7e --- /dev/null +++ b/core/src/test/python/synapsemltest/stages/__init__.py @@ -0,0 +1,2 @@ +# Copyright (C) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. See LICENSE in project root for information. diff --git a/core/src/test/python/synapsemltest/stages/test_ensemble_by_key.py b/core/src/test/python/synapsemltest/stages/test_ensemble_by_key.py new file mode 100644 index 00000000000..ba470bf9cd7 --- /dev/null +++ b/core/src/test/python/synapsemltest/stages/test_ensemble_by_key.py @@ -0,0 +1,57 @@ +# Copyright (C) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. See LICENSE in project root for information. + +import tempfile +import unittest +from pathlib import Path + +from synapse.ml.core.init_spark import init_spark +from synapse.ml.stages import EnsembleByKey + +spark = init_spark() + + +class EnsembleByKeySpec(unittest.TestCase): + def test_col_names_follow_params_after_transform_and_load(self): + frame = spark.createDataFrame( + [("group", 1.0, 2.0), ("group", 3.0, 4.0)], + ["key", "score", "other"], + ) + with self.assertRaisesRegex(Exception, "keys must be set and non-empty"): + EnsembleByKey(keys=[], cols=["score"]).transform(frame) + with self.assertRaisesRegex(Exception, "cols must be set and non-empty"): + EnsembleByKey(keys=["key"], cols=[]).transform(frame) + + transformer = EnsembleByKey(keys=["key"], cols=["score"]) + + self.assertEqual(transformer.getColNames(), ["mean(score)"]) + self.assertFalse(transformer.isSet(transformer.colNames)) + self.assertFalse(transformer.hasDefault(transformer.colNames)) + transformer.transform(frame).collect() + self.assertFalse(transformer.isSet(transformer.colNames)) + self.assertFalse(transformer.hasDefault(transformer.colNames)) + + transformer.setCols(["score", "other"]) + self.assertEqual(transformer.getColNames(), ["mean(score)", "mean(other)"]) + + with tempfile.TemporaryDirectory() as directory: + model_path = str(Path(directory) / "ensemble-by-key") + transformer.write().save(model_path) + loaded = EnsembleByKey.load(model_path) + + self.assertEqual(loaded.getColNames(), ["mean(score)", "mean(other)"]) + self.assertFalse(loaded.isSet(loaded.colNames)) + self.assertFalse(loaded.hasDefault(loaded.colNames)) + + transformer.setColNames(["average-score", "average-other"]) + with tempfile.TemporaryDirectory() as directory: + model_path = str(Path(directory) / "ensemble-by-key-explicit") + transformer.write().save(model_path) + loaded = EnsembleByKey.load(model_path) + + self.assertEqual(loaded.getColNames(), ["average-score", "average-other"]) + self.assertTrue(loaded.isSet(loaded.colNames)) + + +if __name__ == "__main__": + unittest.main() diff --git a/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeyResolutionSuite.scala b/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeyResolutionSuite.scala new file mode 100644 index 00000000000..517b8efbaa8 --- /dev/null +++ b/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeyResolutionSuite.scala @@ -0,0 +1,163 @@ +// Copyright (C) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. See LICENSE in project root for information. + +package com.microsoft.azure.synapse.ml.stages + +import com.microsoft.azure.synapse.ml.core.test.base.TestBase +import org.apache.spark.SparkException +import org.apache.spark.ml.Pipeline +import org.apache.spark.ml.linalg.{SQLDataTypes, Vector} +import org.apache.spark.sql.functions.{col, struct} +import org.apache.spark.sql.types.{DoubleType, IntegerType, Metadata, StringType, StructField, StructType} +import org.apache.spark.sql.{AnalysisException, DataFrame, Row} + +/** Covers the duplicate attribute resolution rules that EnsembleByKey mirrors from Spark's + * `AttributeSeq.resolve`. + */ +class EnsembleByKeyResolutionSuite extends TestBase { + + private val duplicateKey = "__is_duplicate" + + test("custom stage identifiers should not affect internal column resolution") { + val input = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("key", "score") + Seq("ensemble.by.key", "ensemble`by`key").foreach { uid => + val transformer = new EnsembleByKey(uid) + .setKey("key").setCol("score").setCollapseGroup(false) + assert(transformer.transformSchema(input.schema) === transformer.transform(input).schema) + } + } + + test("non-collapsed output should retain rows with null grouping keys") { + val schema = StructType(Array( + StructField("id", IntegerType, nullable = false), + StructField("key", StringType), + StructField("score", DoubleType, nullable = false))) + val missingKey = Option.empty[String].orNull + val input = spark.createDataFrame(java.util.Arrays.asList( + Row(0, missingKey, 1.0), + Row(1, missingKey, 3.0), + Row(2, "group", 5.0)), schema) + val transformed = new EnsembleByKey() + .setKey("key").setCol("score").setCollapseGroup(false).transform(input) + + assert(transformed.orderBy("id").collect().map(row => + (row.getInt(1), Option(row.getString(0)), row.getDouble(3))) === + Array((0, None, 2.0), (1, None, 2.0), (2, Some("group"), 5.0))) + } + + test("vector mean schema should match Spark for all-null inputs") { + val schema = StructType(Array( + StructField("key", StringType, nullable = false), + StructField("features", SQLDataTypes.VectorType))) + val missingVector = Option.empty[Vector].orNull + val input = spark.createDataFrame(java.util.Arrays.asList( + Row("group", missingVector), + Row("group", missingVector)), schema) + val transformer = new EnsembleByKey().setKey("key").setCol("features") + val transformed = transformer.transform(input) + + assert(transformer.transformSchema(input.schema) === transformed.schema) + assert(!transformed.schema("mean(features)").nullable) + intercept[SparkException](transformed.collect()) + } + + test("duplicated qualifier attributes should follow Spark expression identity") { + val base = spark.createDataFrame(Seq(("top", "nested", 1.0), ("top", "nested", 3.0))) + .toDF("group", "nestedGroup", "score") + val nestedGroup = struct(col("nestedGroup").alias("group")).alias("dup") + val shared = base.select(col("group"), col("group"), nestedGroup, col("score")) + val transformer = new EnsembleByKey().setKey("dup.group").setCol("score") + + assert(distinctExpressions(shared, "group") === 1) + Seq("dup" -> "top", "other" -> "nested").foreach { case (alias, expected) => + val transformed = assertSchemaAgrees(transformer, shared.as(alias)) + withClue(s"$alias: ") { + assert(transformed.head().getString(0) === expected) + assert(transformed.select("mean(score)").head().getDouble(0) === 2.0) + } + } + + val ambiguous = base.select(col("group"), nestedGroup, col("score")).as("dup") + .crossJoin(spark.createDataFrame(Seq(Tuple1("side"))).toDF("group").as("dup")) + assert(distinctExpressions(ambiguous, "group") === 2) + intercept[AnalysisException](ambiguous.select("dup.group")) + val error = intercept[IllegalArgumentException](transformer.transform(ambiguous)) + assert(error.getMessage.contains("dup.group is ambiguous")) + } + + test("duplicated unqualified attributes sharing one expression should aggregate") { + val base = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("key", "score") + val duplicated = base.select(col("key"), col("score"), col("score")) + val transformer = new EnsembleByKey().setKey("key").setCol("score") + + assert(duplicated.schema.fieldNames === Array("key", "score", "score")) + assert(distinctExpressions(duplicated, "score") === 1) + assert(duplicated.select("score").columns === Array("score")) + + val transformed = transformer.transform(duplicated) + assert(transformed.schema.fieldNames === Array("key", "mean(score)")) + assert(transformed.schema("mean(score)") === StructField("mean(score)", DoubleType)) + assert(transformed.head().getDouble(1) === 2.0) + + assert(transformer.transformSchema(duplicated.schema) === transformed.schema) + val pipelineModel = new Pipeline().setStages(Array(transformer)).fit(duplicated) + assert(pipelineModel.transform(duplicated).collect() === transformed.collect()) + } + + test("union duplicate attributes should follow Spark duplicate pruning") { + val base = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("key", "score") + val duplicated = base.select(col("key"), col("score"), col("score")) + val unioned = duplicated.union(duplicated) + assert(unioned.schema.fieldNames === Array("key", "score", "score")) + assert(unioned.schema.fields.last.metadata.contains(duplicateKey)) + assert(distinctExpressions(unioned, "score") === 2) + assert(unioned.select("score").columns === Array("score")) + + val transformed = assertSchemaAgrees(new EnsembleByKey().setKey("key").setCol("score"), unioned) + assert(transformed.schema.fieldNames === Array("key", "mean(score)")) + assert(transformed.head().getDouble(1) === 2.0) + + val qualified = assertSchemaAgrees( + new EnsembleByKey().setKey("key").setCol("u.score"), unioned.as("u")) + assert(qualified.schema.fieldNames === Array("key", "mean(u.score)")) + assert(qualified.head().getDouble(1) === 2.0) + } + + test("duplicate pruning should not override qualifier selection") { + // The only `group` attribute of `u` carries Spark's duplicate marker while `v.group` does not, + // so pruning before qualifier selection would silently resolve `u.group` to `v.group`. + val base = spark.createDataFrame(Seq(("u", 1.0), ("u", 3.0))).toDF("group", "score") + val duplicated = base.select(col("group"), col("group"), col("score")) + val tagged = duplicated.union(duplicated).toDF("other", "group", "score") + assert(tagged.schema("group").metadata.contains(duplicateKey)) + assert(!tagged.schema("other").metadata.contains(duplicateKey)) + + val untagged = spark.createDataFrame(Seq(Tuple1("v"))).toDF("group") + val joined = tagged.as("u").crossJoin(untagged.as("v")) + assert(joined.schema.fieldNames === Array("other", "group", "score", "group")) + assert(joined.select("u.group").head().getString(0) === "u") + + val transformed = assertSchemaAgrees( + new EnsembleByKey().setKey("u.group").setCol("score"), joined) + assert(transformed.schema.fieldNames === Array("group", "mean(score)")) + assert(transformed.schema("group").metadata === Metadata.empty) + assert(transformed.head().getString(0) === "u") + assert(transformed.head().getDouble(1) === 2.0) + + val nonCollapsed = new EnsembleByKey().setKey("u.group").setCol("score").setCollapseGroup(false) + val schemaError = intercept[IllegalArgumentException](nonCollapsed.transformSchema(joined.schema)) + val transformError = intercept[IllegalArgumentException](nonCollapsed.transform(joined)) + assert(schemaError.getMessage.contains("multiple columns are named group")) + assert(transformError.getMessage.contains("multiple columns are named group")) + } + + private def distinctExpressions(input: DataFrame, name: String): Int = { + input.queryExecution.analyzed.output.filter(_.name == name).map(_.exprId).distinct.length + } + + private def assertSchemaAgrees(transformer: EnsembleByKey, input: DataFrame): DataFrame = { + val transformed = transformer.transform(input) + assert(transformer.transformSchema(input.schema) === transformed.schema) + transformed + } +} diff --git a/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeySuite.scala b/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeySuite.scala index 1a624cf4431..9fd5c993757 100644 --- a/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeySuite.scala +++ b/core/src/test/scala/com/microsoft/azure/synapse/ml/stages/EnsembleByKeySuite.scala @@ -5,9 +5,14 @@ package com.microsoft.azure.synapse.ml.stages import com.microsoft.azure.synapse.ml.core.test.base.TestBase import com.microsoft.azure.synapse.ml.core.test.fuzzing.{TestObject, TransformerFuzzing} +import org.apache.spark.ml.Pipeline import org.apache.spark.ml.feature.VectorAssembler -import org.apache.spark.ml.linalg.DenseVector -import org.apache.spark.sql.DataFrame +import org.apache.spark.ml.linalg.{DenseVector, SQLDataTypes} +import org.apache.spark.sql.{AnalysisException, DataFrame, Row, SparkSession} +import org.apache.spark.sql.catalyst.expressions.{Cast, RowOrdering} +import org.apache.spark.sql.functions.{array, col, expr, lit, map, struct} +import org.apache.spark.sql.types.{CalendarIntervalType, DoubleType, MapType, Metadata, StringType, + StructField, StructType} class EnsembleByKeySuite extends TestBase with TransformerFuzzing[EnsembleByKey] { @@ -53,6 +58,623 @@ class EnsembleByKeySuite extends TestBase with TransformerFuzzing[EnsembleByKey] df1.show() } + test("transformSchema should match mixed aggregate output for default and explicit names") { + val input = mixedTypeDF + val inputNames = Array("doubleScore", "floatScore", "features") + val defaultNames = inputNames.map(name => s"mean($name)") + val explicitNames = Array("averageDouble", "averageFloat", "averageFeatures") + val keyNames = Array("group", "region") + + assert(input.schema("features").metadata !== Metadata.empty) + + Seq(defaultNames -> false, explicitNames -> true).foreach { case (outputNames, useExplicitNames) => + Seq(true, false).foreach { collapseGroup => + val transformer = new EnsembleByKey() + .setKeys(keyNames).setCols(inputNames).setCollapseGroup(collapseGroup) + if (useExplicitNames) { + transformer.setColNames(outputNames) + } + + val transformedSchema = transformer.transformSchema(input.schema) + val actualSchema = transformer.transform(input).schema + val expectedNames = if (collapseGroup) { + keyNames ++ outputNames + } else { + keyNames ++ input.columns.filterNot((keyNames ++ outputNames).contains) ++ outputNames + } + + withClue(s"explicitNames=$useExplicitNames, collapseGroup=$collapseGroup: ") { + assert(transformedSchema === actualSchema) + assert(actualSchema.fieldNames === expectedNames) + assert(actualSchema(outputNames(0)) === StructField(outputNames(0), DoubleType)) + assert(actualSchema(outputNames(1)) === StructField(outputNames(1), DoubleType)) + assert(actualSchema(outputNames(2)) === + StructField(outputNames(2), SQLDataTypes.VectorType, nullable = false)) + } + } + } + } + + test("non-collapsed output should overwrite numeric and vector columns") { + val input = mixedTypeDF + val overwrittenNames = Array("doubleScore", "floatScore", "features") + val transformer = new EnsembleByKey() + .setKeys("group", "region").setCols(overwrittenNames) + .setColNames(overwrittenNames).setCollapseGroup(false) + + val transformedSchema = transformer.transformSchema(input.schema) + val transformed = transformer.transform(input) + + assert(transformed.schema === transformedSchema) + assert(transformed.columns === + Array("group", "region", "id", "component1", "component2") ++ overwrittenNames) + assert(transformed.schema("features").metadata === Metadata.empty) + assert(!transformed.schema("features").nullable) + + val actual = transformed.orderBy("id") + .select("doubleScore", "floatScore", "features") + .collect() + .map(row => (row.getDouble(0), row.getDouble(1), row.getAs[DenseVector](2))) + val expected = Array( + (1.0, 1.0, new DenseVector(Array(1.0, 0.1))), + (2.0, 2.0, new DenseVector(Array(2.0, -2.5))), + (2.0, 2.0, new DenseVector(Array(2.0, -2.5)))) + + assert(actual === expected) + } + + test("non-collapsed output should replace case-variant columns consistently") { + val input = spark.createDataFrame(Seq((0, "group", 1.0, "lower", "upper"))) + .toDF("id", "key", "score", "features", "FEATURES") + + Seq(false -> Array("key", "id", "score", "features"), + true -> Array("key", "id", "score", "FEATURES", "features")) + .foreach { case (caseSensitive, expectedNames) => + withCaseSensitiveAnalysis(caseSensitive) { + val transformer = new EnsembleByKey() + .setKey("key").setCol("score").setColName("features").setCollapseGroup(false) + + val transformedSchema = transformer.transformSchema(input.schema) + val actualSchema = transformer.transform(input).schema + + assert(transformedSchema === actualSchema) + assert(actualSchema.fieldNames === expectedNames) + } + } + } + + test("default output names should follow updated input columns before transform") { + val transformer = new EnsembleByKey().setKeys("group", "region").setCol("doubleScore") + assert(transformer.getDefault(transformer.colNames).isEmpty) + transformer.transformSchema(mixedTypeDF.schema) + assert(transformer.getDefault(transformer.colNames).isEmpty) + assert(transformer.getColNames === Array("mean(doubleScore)")) + transformer.transform(mixedTypeDF) + assert(transformer.getDefault(transformer.colNames).isEmpty) + transformer.setCols("doubleScore", "floatScore") + + assert(transformer.transformSchema(mixedTypeDF.schema).fieldNames === + Array("group", "region", "mean(doubleScore)", "mean(floatScore)")) + assert(transformer.getColNames === Array("mean(doubleScore)", "mean(floatScore)")) + } + + test("grouping keys should resolve case-insensitively to input field names") { + withCaseSensitiveAnalysis(false) { + Seq(true, false).foreach { collapseGroup => + val transformer = new EnsembleByKey() + .setKeys("GROUP", "REGION").setCol("doubleScore").setCollapseGroup(collapseGroup) + + val transformedSchema = transformer.transformSchema(mixedTypeDF.schema) + val actualSchema = transformer.transform(mixedTypeDF).schema + + withClue(s"collapseGroup=$collapseGroup: ") { + assert(transformedSchema === actualSchema) + assert(actualSchema.fieldNames.take(2) === Array("GROUP", "REGION")) + } + } + } + } + + test("grouping key resolution should honor case-sensitive analysis") { + withCaseSensitiveAnalysis(true) { + val input = spark.createDataFrame(Seq(("lower", "upper", 1.0))) + .toDF("group", "GROUP", "score") + val transformer = new EnsembleByKey().setKey("group").setCol("score") + + assert(transformer.transformSchema(input.schema) === transformer.transform(input).schema) + + val error = intercept[IllegalArgumentException] { + new EnsembleByKey().setKey("Group").setCol("score").transformSchema(input.schema) + } + assert(error.getMessage.contains("Group does not exist")) + } + } + + test("transformSchema should match output when grouping column retention is disabled") { + withSQLConf("spark.sql.retainGroupColumns", "false") { + Seq(true, false).foreach { collapseGroup => + val transformer = new EnsembleByKey() + .setKeys("group", "region").setCol("doubleScore").setCollapseGroup(collapseGroup) + val transformedSchema = transformer.transformSchema(mixedTypeDF.schema) + val actualSchema = transformer.transform(mixedTypeDF).schema + + assert(transformedSchema === actualSchema) + assert(actualSchema.fieldNames.take(2) === Array("group", "region")) + } + } + } + + test("transform should use the dataset session for grouping column retention") { + val disabledSession = spark.newSession() + disabledSession.conf.set("spark.sql.retainGroupColumns", false) + val disabledInput = disabledSession.createDataFrame(Seq(("group", 1.0))).toDF("group", "score") + + withActiveSession(spark) { + val transformer = new EnsembleByKey().setKey("group").setCol("score") + assert(transformer.transformSchema(disabledInput.schema) === transformer.transform(disabledInput).schema) + } + + val enabledSession = spark.newSession() + enabledSession.conf.set("spark.sql.retainGroupColumns", true) + val enabledInput = enabledSession.createDataFrame(Seq(("group", 1.0))).toDF("group", "score") + + withSQLConf("spark.sql.retainGroupColumns", "false") { + withActiveSession(spark) { + val transformer = new EnsembleByKey().setKey("group").setCol("score") + val transformed = transformer.transform(enabledInput) + val pipelineModel = new Pipeline().setStages(Array(transformer)).fit(enabledInput) + + assert(transformer.transformSchema(enabledInput.schema) === transformed.schema) + assert(pipelineModel.transform(enabledInput).schema === transformed.schema) + assert(transformed.columns === Array("group", "mean(score)")) + } + } + } + + test("configuration parsing should match Spark boolean parsing") { + withSQLConf("spark.sql.caseSensitive", " false ") { + val transformer = new EnsembleByKey().setKey("GROUP").setCol("doubleScore") + assert(transformer.transformSchema(mixedTypeDF.schema) === transformer.transform(mixedTypeDF).schema) + } + + withSQLConf("spark.sql.retainGroupColumns", " true ") { + val transformer = new EnsembleByKey().setKey("group").setCol("doubleScore") + assert(transformer.transformSchema(mixedTypeDF.schema) === transformer.transform(mixedTypeDF).schema) + } + + withSQLConf("spark.sql.retainGroupColumns", " false ") { + val transformer = new EnsembleByKey().setKey("group").setCol("doubleScore") + assert(transformer.transformSchema(mixedTypeDF.schema) === transformer.transform(mixedTypeDF).schema) + } + } + + test("no active session should expose the documented case-resolution limitation") { + withSQLConf("spark.sql.caseSensitive", "true") { + val input = spark.createDataFrame(Seq((0, "group", 1.0, 2.0, 3.0))) + .toDF("id", "key", "score", "features", "FEATURES") + val transformer = new EnsembleByKey() + .setKey("key").setCol("score").setColName("features").setCollapseGroup(false) + val assembler = new VectorAssembler() + .setInputCols(Array("FEATURES")).setOutputCol("vector") + val pipeline = new Pipeline().setStages(Array(transformer, assembler)) + + withoutActiveSession { + val transformedSchema = transformer.transformSchema(input.schema) + val actualSchema = transformer.transform(input).schema + + assert(transformedSchema.fieldNames === Array("key", "id", "score", "features")) + assert(actualSchema.fieldNames === Array("key", "id", "score", "FEATURES", "features")) + val pipelineError = intercept[IllegalArgumentException](pipeline.fit(input)) + assert(pipelineError.getMessage.contains("FEATURES does not exist")) + } + pipeline.fit(input) + } + } + + test("transform should use the dataset session for column resolution") { + val sensitiveSession = spark.newSession() + sensitiveSession.conf.set("spark.sql.caseSensitive", true) + val sensitiveInput = sensitiveSession.createDataFrame(Seq(("group", 1.0))).toDF("group", "score") + + withCaseSensitiveAnalysis(false) { + val transformer = new EnsembleByKey().setKey("GROUP").setCol("SCORE") + assert(transformer.transformSchema(sensitiveInput.schema).fieldNames === Array("GROUP", "mean(SCORE)")) + assert(intercept[IllegalArgumentException](transformer.transform(sensitiveInput)) + .getMessage.contains("does not exist")) + } + + val insensitiveSession = spark.newSession() + insensitiveSession.conf.set("spark.sql.caseSensitive", false) + val insensitiveInput = insensitiveSession.createDataFrame(Seq(("group", 1.0))).toDF("group", "score") + + withCaseSensitiveAnalysis(true) { + val transformer = new EnsembleByKey().setKey("GROUP").setCol("SCORE") + assert(transformer.transform(insensitiveInput).schema.fieldNames === Array("GROUP", "mean(SCORE)")) + } + + withoutActiveSession { + val transformer = new EnsembleByKey().setKey("GROUP").setCol("SCORE") + assert(intercept[IllegalArgumentException](transformer.transform(sensitiveInput)) + .getMessage.contains("does not exist")) + } + } + + test("nested and quoted field references should match Spark resolution") { + val nestedInput = spark.createDataFrame(Seq(("a", 1.0), ("a", 3.0))) + .toDF("nestedKey", "score") + .select(struct(col("nestedKey").alias("key")).alias("nested"), col("score")) + + Seq(true, false).foreach { collapseGroup => + val transformer = new EnsembleByKey() + .setKey("nested.key").setCol("score").setCollapseGroup(collapseGroup) + val transformedSchema = transformer.transformSchema(nestedInput.schema) + val transformed = transformer.transform(nestedInput) + + assert(transformedSchema === transformed.schema) + assert(transformed.schema.fieldNames.head === "key") + assert(transformed.select("mean(score)").head().getDouble(0) === 2.0) + } + + val dottedInput = spark.createDataFrame(Seq(("a", 1.0), ("a", 3.0))).toDF("a.b", "score") + val dottedTransformer = new EnsembleByKey().setKey("`a.b`").setCol("score") + + assert(dottedTransformer.transformSchema(dottedInput.schema) === dottedTransformer.transform(dottedInput).schema) + } + + test("nested key nullability should include nullable ancestor structs") { + val inputSchema = StructType(Array( + StructField( + "nested", + StructType(Array(StructField("key", StringType, nullable = false))), + nullable = true), + StructField("score", DoubleType, nullable = false))) + val rows = java.util.Arrays.asList( + Row(Row("a"), 1.0), + Row(Row("a"), 3.0)) + val input = spark.createDataFrame(rows, inputSchema) + + Seq(true, false).foreach { collapseGroup => + val transformer = new EnsembleByKey() + .setKey("nested.key").setCol("score").setCollapseGroup(collapseGroup) + val transformedSchema = transformer.transformSchema(input.schema) + val actualSchema = transformer.transform(input).schema + + assert(transformedSchema === actualSchema) + assert(actualSchema("key").nullable) + } + } + + test("non-collapsed nested keys should reject unsafe leaf-name collisions") { + val collisionInput = spark.createDataFrame(Seq(("row-1", "group", 1.0))) + .toDF("id", "nestedId", "score") + .select(col("id"), struct(col("nestedId").alias("id")).alias("meta"), col("score")) + val collisionTransformer = new EnsembleByKey() + .setKey("meta.id").setCol("score").setCollapseGroup(false) + + assertConsistentSchemaError( + collisionTransformer, collisionInput, "ambiguous between a nested field and a dataset qualifier") + + val duplicateInput = spark.createDataFrame(Seq(("left", "right", 1.0))) + .toDF("leftKey", "rightKey", "score") + .select( + struct(col("leftKey").alias("key")).alias("left"), + struct(col("rightKey").alias("key")).alias("right"), + col("score")) + val duplicateTransformer = new EnsembleByKey() + .setKeys("left.key", "right.key").setCol("score").setCollapseGroup(false) + + assertConsistentSchemaError(duplicateTransformer, duplicateInput, "must resolve to distinct output columns") + } + + test("non-collapsed duplicate grouping keys should fail consistently") { + Seq("true", "false").foreach { retainGroupColumns => + withSQLConf("spark.sql.retainGroupColumns", retainGroupColumns) { + val transformer = new EnsembleByKey() + .setKeys("group", "group").setCol("doubleScore").setCollapseGroup(false) + + assertConsistentSchemaError( + transformer, + mixedTypeDF, + "must resolve to distinct output columns") + } + } + + val collapsed = new EnsembleByKey() + .setKeys("group", "group").setCol("doubleScore").setCollapseGroup(true) + assert(collapsed.transformSchema(mixedTypeDF.schema) === collapsed.transform(mixedTypeDF).schema) + } + + test("nested keys should preserve unreferenced duplicate top-level columns") { + val input = spark.createDataFrame(Seq(("group", 1.0, 10.0))) + .toDF("key", "score", "duplicate") + .select( + struct(col("key").alias("value")).alias("nested"), + col("score"), + col("duplicate").alias("duplicate"), + col("duplicate").alias("duplicate")) + val transformer = new EnsembleByKey() + .setKey("nested.value").setCol("score").setCollapseGroup(false) + + val transformed = transformer.transform(input) + assert(transformer.transformSchema(input.schema) === transformed.schema) + assert(transformed.schema.fieldNames === + Array("value", "nested", "score", "duplicate", "duplicate", "mean(score)")) + } + + test("quoted field references should ignore quoted-regex column settings") { + withSQLConf("spark.sql.parser.quotedRegexColumnNames", "true") { + val keyInput = spark.createDataFrame(Seq(("a", 1.0), ("a", 3.0))).toDF("a.b", "score") + val keyTransformer = new EnsembleByKey().setKey("`a.b`").setCol("score") + assert(keyTransformer.transformSchema(keyInput.schema) === keyTransformer.transform(keyInput).schema) + + val colInput = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("group", "s.c") + val colTransformer = new EnsembleByKey().setKey("group").setCol("`s.c`") + assert(colTransformer.transformSchema(colInput.schema) === colTransformer.transform(colInput).schema) + } + } + + test("literal dotted aggregate columns should require Spark quoting") { + val input = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("group", "s.c") + val quotedTransformer = new EnsembleByKey().setKey("group").setCol("`s.c`") + + assert(quotedTransformer.transformSchema(input.schema) === quotedTransformer.transform(input).schema) + + val plainTransformer = new EnsembleByKey().setKey("group").setCol("s.c") + assertConsistentSchemaError(plainTransformer, input, "s.c does not exist") + } + + test("qualified and collection field references should match Spark resolution") { + val qualifiedInput = mixedTypeDF.as("source") + val qualifiedTransformer = new EnsembleByKey().setKey("source.group").setCol("doubleScore") + assert(qualifiedTransformer.transformSchema(qualifiedInput.schema) === + qualifiedTransformer.transform(qualifiedInput).schema) + + val collectionBase = spark.createDataFrame(Seq(("group", 1.0))).toDF("key", "score") + val arrayInput = collectionBase.select( + array(struct(col("key").alias("field"))).alias("items"), + col("score")) + val arrayTransformer = new EnsembleByKey().setKey("items.field").setCol("score") + val arrayResult = arrayTransformer.transform(arrayInput) + assert(arrayTransformer.transformSchema(arrayInput.schema) === arrayResult.schema) + assert(arrayResult.collect().head.getSeq[String](0) === Seq("group")) + + val nullableArrayInput = collectionBase.select( + array(struct(expr("CAST(NULL AS STRING)").alias("field"))).alias("items"), + col("score")) + val nullableArrayTransformer = new EnsembleByKey().setKey("items.field").setCol("score") + val nullableArrayResult = nullableArrayTransformer.transform(nullableArrayInput) + assert(nullableArrayTransformer.transformSchema(nullableArrayInput.schema) === nullableArrayResult.schema) + assert(Option(nullableArrayResult.collect().head.getSeq[String](0).head).isEmpty) + + val mapInput = collectionBase.select( + map(lit("field"), col("key")).alias("values"), + col("score")) + val mapTransformer = new EnsembleByKey().setKey("values.field").setCol("score") + assert(mapTransformer.transformSchema(mapInput.schema) === mapTransformer.transform(mapInput).schema) + + val invalidMapInput = collectionBase.select( + map(struct(lit(1).alias("part")), col("key")).alias("values"), + col("score")) + val invalidMapTransformer = new EnsembleByKey().setKey("values.field").setCol("score") + assertConsistentSchemaError(invalidMapTransformer, invalidMapInput, "does not accept string keys") + } + + test("map key extraction should follow Spark cast coercion") { + val base = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("key", "score") + Seq( + "values.true" -> map(lit(true), col("key")), + "values.field" -> map(lit("field").cast("binary"), col("key")), + "values.1" -> map(lit(1), col("key")), + "values.2020-01-01" -> map(lit("2020-01-01").cast("date"), col("key")) + ).foreach { case (reference, values) => + val input = base.select(values.alias("values"), col("score")) + val transformed = assertSchemaAgrees(new EnsembleByKey().setKey(reference).setCol("score"), input) + withClue(s"$reference: ") { + assert(transformed.head().getString(0) === "group") + assert(transformed.select("mean(score)").head().getDouble(0) === 2.0) + } + } + } + + test("map keys Spark cannot order should be rejected consistently") { + val base = spark.createDataFrame(Seq(("group", 1.0), ("group", 3.0))).toDF("key", "score") + val input = base.select( + map(expr("make_interval(0, 0, 0, 1, 0, 0, 0)"), col("key")).alias("values"), + col("score")) + val keyType = input.schema("values").dataType.asInstanceOf[MapType].keyType + + assert(keyType === CalendarIntervalType) + assert(Cast.canCast(StringType, keyType), "the key type is castable from a string literal") + assert(!RowOrdering.isOrderable(keyType), "the key type is not orderable, so GetMapValue fails") + intercept[AnalysisException](input.select(expr("values[make_interval(0, 0, 0, 1, 0, 0, 0)]")).schema) + + val transformer = new EnsembleByKey().setKey("values.1 days").setCol("score") + assertConsistentSchemaError(transformer, input, "map key type CalendarIntervalType is not orderable") + assertConsistentSchemaError(transformer, input, "Use a map column whose key type is orderable") + } + + test("extracted grouping values Spark cannot order should be rejected consistently") { + val input = spark.range(1).select( + map(lit("outer"), map(lit("inner"), lit(1))).alias("values"), + lit(1.0).alias("score")) + val transformer = new EnsembleByKey().setKey("values.outer").setCol("score") + + assertConsistentSchemaError(transformer, input, "Spark cannot use as a grouping key") + } + + test("map extraction should reject dataset qualifier collisions") { + val input = spark.createDataFrame(Seq(("group", 1.0))) + .toDF("key", "score") + .select(map(lit("field"), col("score")).alias("values"), col("score"), col("key").alias("field")) + .as("values") + val transformer = new EnsembleByKey().setKey("values.field").setCol("score") + + assertConsistentSchemaError(transformer, input, "ambiguous between a nested field and a dataset qualifier") + } + + test("nested key output names should preserve configured casing") { + val input = spark.createDataFrame(Seq(("group", 1.0))) + .toDF("key", "score") + .select(struct(col("key").alias("Key")).alias("nested"), col("score")) + val transformer = new EnsembleByKey().setKey("nested.key").setCol("score") + + assert(transformer.transformSchema(input.schema) === transformer.transform(input).schema) + assert(transformer.transform(input).schema.fieldNames.head === "key") + } + + test("qualified references should preserve qualifier identity") { + val left = spark.createDataFrame(Seq((1, "left", 1.0))).toDF("id", "group", "score").as("left") + val right = spark.createDataFrame(Seq((1, "right"))).toDF("id", "group").as("right") + val joined = left.join(right, Seq("id")) + + Seq("left", "right").foreach { qualifier => + val transformer = new EnsembleByKey().setKey(s"$qualifier.group").setCol("score") + withClue(s"$qualifier: ") { + assert(assertSchemaAgrees(transformer, joined).head().getString(0) === qualifier) + } + } + + assertConsistentSchemaError( + new EnsembleByKey().setKey("right.group").setCol("score").setCollapseGroup(false), + joined, + "multiple columns are named group when collapseGroup is false") + + val invalidQualifier = new EnsembleByKey().setKey("wrong.group").setCol("score") + assert(invalidQualifier.transformSchema(joined.schema).fieldNames === Array("group", "mean(score)")) + val error = intercept[IllegalArgumentException](invalidQualifier.transform(joined)) + assert(error.getMessage.contains("does not match a dataset qualifier")) + } + + test("non-collapsed qualified references should preserve unrelated duplicates") { + val left = spark.createDataFrame(Seq((1, "group", 1.0))).toDF("id", "group", "score").as("left") + val right = spark.createDataFrame(Seq((1, 2.0))).toDF("id", "score").as("right") + val joined = left.join(right, Seq("id")) + val transformer = new EnsembleByKey() + .setKey("left.group").setCol("left.score") + .setColName("average").setCollapseGroup(false) + val transformed = assertSchemaAgrees(transformer, joined) + + assert(transformed.schema.fieldNames === Array("group", "id", "score", "score", "average")) + assert(transformed.head().getDouble(4) === 1.0) + } + + test("qualified aggregates should compare derived aggregate outputs") { + val left = spark.createDataFrame(Seq((1, "group", 1.0), (2, "group", 3.0))) + .toDF("id", "group", "score").as("left") + val right = spark.createDataFrame(Seq((1, 5.0))).toDF("id", "score").as("right") + val joined = left.join(right, Seq("id"), "left_outer") + assert(joined.schema.fields.filter(_.name == "score").map(_.nullable) === Array(false, true)) + + Seq("left.score" -> 2.0, "right.score" -> 5.0).foreach { case (reference, expected) => + val transformed = assertSchemaAgrees(new EnsembleByKey().setKey("group").setCol(reference), joined) + withClue(s"$reference: ") { + assert(transformed.schema.last === StructField(s"mean($reference)", DoubleType)) + assert(transformed.head().getDouble(1) === expected) + } + } + + assertConsistentSchemaError( + new EnsembleByKey().setKey("right.score").setCol("left.score"), + joined, + "incompatible declared outputs") + + val nestedLeft = spark.createDataFrame(Seq((1, 1.0), (1, 3.0))).toDF("id", "value") + .select(col("id"), struct(col("value")).alias("s")).as("left") + val nestedRight = spark.createDataFrame(Seq((1, 5.0f))).toDF("id", "value") + .select(col("id"), struct(col("value")).alias("s")).as("right") + val nested = assertSchemaAgrees( + new EnsembleByKey().setKey("id").setCol("right.s.value"), + nestedLeft.join(nestedRight, Seq("id"))) + assert(nested.schema.last === StructField("mean(right.s.value)", DoubleType)) + assert(nested.head().getDouble(1) === 5.0) + + val stringRight = spark.createDataFrame(Seq((1, "5"))).toDF("id", "value") + .select(col("id"), struct(col("value")).alias("s")).as("right") + assertConsistentSchemaError( + new EnsembleByKey().setKey("id").setCol("right.s.value"), + nestedLeft.join(stringRight, Seq("id")), + "incompatible declared outputs") + } + + test("multipart qualifiers should agree with schema-only interpretations") { + val base = spark.createDataFrame(Seq(("top", "nested", 1.0), ("top", "nested", 3.0))) + .toDF("group", "nestedGroup", "score") + val viewName = s"ensembleView${System.nanoTime()}" + val input = base.select( + col("group"), struct(col("nestedGroup").alias("group")).alias(viewName), col("score")) + val transformer = new EnsembleByKey().setKey(s"global_temp.$viewName.group").setCol("score") + + assert(assertSchemaAgrees(transformer, input.as("global_temp")).head().getString(0) === "nested") + + input.createOrReplaceGlobalTempView(viewName) + try { + val view = spark.table(s"global_temp.$viewName") + assert(assertSchemaAgrees(transformer, view).head().getString(0) === "top") + } finally { + spark.catalog.dropGlobalTempView(viewName) + } + + val conflicting = base.select( + col("score").alias("group"), + struct(col("nestedGroup").alias("group")).alias("view"), + col("score")) + assertConsistentSchemaError( + new EnsembleByKey().setKey("global_temp.view.group").setCol("score"), + conflicting, + "ambiguous between a nested field and a dataset qualifier") + } + + test("schema and runtime should reject invalid column configurations") { + val invalidConfigurations = Seq( + new EnsembleByKey().setCol("doubleScore") -> "keys must be set and non-empty", + new EnsembleByKey().setKeys(Array.empty[String]).setCol("doubleScore") -> + "keys must be set and non-empty", + new EnsembleByKey().setKey("group") -> "cols must be set and non-empty", + new EnsembleByKey().setKey("group").setCols(Array.empty[String]) -> + "cols must be set and non-empty", + new EnsembleByKey().setKey("missingKey").setCol("doubleScore") -> "missingKey does not exist", + new EnsembleByKey().setKey("group").setCol("missingCol") -> "missingCol does not exist", + new EnsembleByKey().setKey("group").setCols("doubleScore", "floatScore") + .setColName("average") -> "must have the same length", + new EnsembleByKey().setKey("group").setCol("doubleScore").setColName("GROUP") + .setCollapseGroup(false) -> "cannot overwrite grouping keys" + ) + + invalidConfigurations.foreach { case (transformer, expectedMessage) => + assertConsistentSchemaError(transformer, mixedTypeDF, expectedMessage) + } + + withCaseSensitiveAnalysis(false) { + val ambiguousInput = spark.createDataFrame(Seq(("lower", "upper", 1.0))) + .toDF("group", "GROUP", "score") + val keyTransformer = new EnsembleByKey().setKey("group").setCol("score") + assert(keyTransformer.transformSchema(ambiguousInput.schema).fieldNames === + Array("group", "mean(score)")) + val error = intercept[IllegalArgumentException](keyTransformer.transform(ambiguousInput)) + assert(error.getMessage.contains("group is ambiguous")) + + val ambiguousAggregateInput = spark.createDataFrame(Seq(("group", 1.0, 2.0))) + .toDF("group", "score", "SCORE") + val aggregateTransformer = new EnsembleByKey().setKey("group").setCol("score") + assert(aggregateTransformer.transformSchema(ambiguousAggregateInput.schema).fieldNames === + Array("group", "mean(score)")) + val aggregateError = + intercept[IllegalArgumentException](aggregateTransformer.transform(ambiguousAggregateInput)) + assert(aggregateError.getMessage.contains("score is ambiguous")) + } + } + + test("transformSchema should reject unsupported aggregate types") { + val input = spark.createDataFrame(Seq(("foo", 1))).toDF("group", "score") + val transformer = new EnsembleByKey().setKey("group").setCol("score") + + val error = intercept[IllegalArgumentException] { + transformer.transformSchema(input.schema) + } + + assert(error.getMessage === "Cannot operate on type IntegerType with strategy mean") + } + lazy val testDF: DataFrame = { val initialTestDF = spark.createDataFrame( Seq((0, "foo", 1.0, .1), @@ -64,6 +686,19 @@ class EnsembleByKeySuite extends TestBase with TransformerFuzzing[EnsembleByKey] .setOutputCol("v1").transform(initialTestDF) } + lazy val mixedTypeDF: DataFrame = { + val initialTestDF = spark.createDataFrame( + Seq((0, "west", "foo", 1.0, 1.0f, 1.0, 0.1), + (1, "east", "bar", 4.0, 4.0f, 4.0, -2.0), + (2, "east", "bar", 0.0, 0.0f, 0.0, -3.0))) + .toDF("id", "region", "group", "doubleScore", "floatScore", "component1", "component2") + + new VectorAssembler() + .setInputCols(Array("component1", "component2")) + .setOutputCol("features") + .transform(initialTestDF) + } + lazy val testModel: EnsembleByKey = new EnsembleByKey().setKey("label1").setCol("score1") .setCollapseGroup(false).setVectorDims(Map("v1"->2)) @@ -98,4 +733,45 @@ class EnsembleByKeySuite extends TestBase with TransformerFuzzing[EnsembleByKey] def testObjects(): Seq[TestObject[EnsembleByKey]] = Seq(new TestObject(testModel, testDF)) def reader: EnsembleByKey.type = EnsembleByKey + + private def withCaseSensitiveAnalysis[T](value: Boolean)(action: => T): T = { + withSQLConf("spark.sql.caseSensitive", value.toString)(action) + } + + private def withSQLConf[T](configName: String, value: String)(action: => T): T = { + val previousValue = spark.conf.get(configName) + spark.conf.set(configName, value) + try action finally spark.conf.set(configName, previousValue) + } + + private def withActiveSession[T](session: SparkSession)(action: => T): T = { + val previousSession = SparkSession.getActiveSession + SparkSession.setActiveSession(session) + try action finally { + previousSession.fold(SparkSession.clearActiveSession())(SparkSession.setActiveSession) + } + } + + private def withoutActiveSession[T](action: => T): T = { + val previousSession = SparkSession.getActiveSession + SparkSession.clearActiveSession() + try action finally previousSession.foreach(SparkSession.setActiveSession) + } + + private def assertSchemaAgrees(transformer: EnsembleByKey, input: DataFrame): DataFrame = { + val transformed = transformer.transform(input) + assert(transformer.transformSchema(input.schema) === transformed.schema) + transformed + } + + private def assertConsistentSchemaError( + transformer: EnsembleByKey, + input: DataFrame, + expectedMessage: String + ): Unit = { + val schemaError = intercept[IllegalArgumentException](transformer.transformSchema(input.schema)) + val transformError = intercept[IllegalArgumentException](transformer.transform(input)) + assert(schemaError.getMessage.contains(expectedMessage)) + assert(transformError.getMessage.contains(expectedMessage)) + } }