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..e0b30b041ed 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 @@ -83,12 +83,15 @@ class EnsembleByKey(val uid: String) extends Transformer setDefault(collapseGroup -> true) + private def setDefaultColNames(): Unit = { + if (get(colNames).isEmpty) { + setDefault(colNames -> getCols.map(name => s"$getStrategy($name)")) + } + } + override def transform(dataset: Dataset[_]): DataFrame = { logTransform[DataFrame]({ - - if (get(colNames).isEmpty) { - setDefault(colNames -> getCols.map(name => s"$getStrategy($name)")) - } + setDefaultColNames() transformSchema(dataset.schema) @@ -130,25 +133,31 @@ class EnsembleByKey(val uid: String) extends Transformer } 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") - } + setDefaultColNames() + + val inputNames = getCols + val outputNames = getColNames + val keyNames = getKeys + + val aggregateFields = inputNames.zip(outputNames).map { case (inputName, outputName) => + val inputField = schema(inputName) + inputField.dataType match { + case _: DoubleType => StructField(outputName, DoubleType) + case _: FloatType => StructField(outputName, DoubleType) + case fdt if fdt == VectorType => StructField(outputName, inputField.dataType) + case t => throw new IllegalArgumentException(s"Cannot operate on type $t with strategy $getStrategy") } } - val keyFields = schema.fields.filter(f => colSet(f.name)) - val fields = - (if (getCollapseGroup) schema.fields else keyFields).++(newFields) + val keyFields = keyNames.map(schema(_)) + val fields = if (getCollapseGroup) { + keyFields ++ aggregateFields + } else { + val keyNameSet = keyNames.toSet + val outputNameSet = outputNames.toSet + val inputFields = schema.fields.filterNot(f => keyNameSet(f.name) || outputNameSet(f.name)) + keyFields ++ inputFields ++ aggregateFields + } new StructType(fields) } 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..a2ac7d36002 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 @@ -53,6 +53,34 @@ class EnsembleByKeySuite extends TestBase with TransformerFuzzing[EnsembleByKey] df1.show() } + test("transformSchema should match the transformed schema for both collapse modes") { + val scoreDFDouble = spark.createDataFrame( + Seq((0, "foo", 1.0), + (1, "bar", 4.0), + (1, "bar", 0.0))) + .toDF("id", "group", "score") + val scoreDFFloat = spark.createDataFrame( + Seq((0, "foo", 1.0f), + (1, "bar", 4.0f), + (1, "bar", 0.0f))) + .toDF("id", "group", "score") + + Seq(scoreDFDouble, scoreDFFloat).foreach { scoreDF => + Seq(true, false).foreach { collapseGroup => + val transformer = new EnsembleByKey() + .setKey("group") + .setCol("score") + .setCollapseGroup(collapseGroup) + val transformedSchema = transformer.transformSchema(scoreDF.schema) + val transformed = transformer.transform(scoreDF) + + withClue(s"dataType=${scoreDF.schema("score").dataType}, collapseGroup=$collapseGroup: ") { + assert(transformed.schema === transformedSchema) + } + } + } + } + lazy val testDF: DataFrame = { val initialTestDF = spark.createDataFrame( Seq((0, "foo", 1.0, .1),