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..e8b12f4a8c0 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 (!isSet(colNames)) { + 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, VectorType, nullable = false) + 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..de0c0ddbd89 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 @@ -6,8 +6,9 @@ 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.feature.VectorAssembler -import org.apache.spark.ml.linalg.DenseVector +import org.apache.spark.ml.linalg.{DenseVector, SQLDataTypes} import org.apache.spark.sql.DataFrame +import org.apache.spark.sql.types.{DoubleType, Metadata, StructField} class EnsembleByKeySuite extends TestBase with TransformerFuzzing[EnsembleByKey] { @@ -53,6 +54,98 @@ 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("default output names should follow updated input columns before transform") { + val transformer = new EnsembleByKey() + .setKeys("group", "region") + .setCol("doubleScore") + + transformer.transformSchema(mixedTypeDF.schema) + transformer.setCols("doubleScore", "floatScore") + + assert(transformer.transformSchema(mixedTypeDF.schema).fieldNames === + Array("group", "region", "mean(doubleScore)", "mean(floatScore)")) + } + + 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 +157,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))