Repository: spark
Updated Branches:
  refs/heads/branch-1.5 5f58704c9 -> 5b7067c91


[SPARK-10573] [ML] IndexToString output schema should be StringType

Fixes bug where IndexToString output schema was DoubleType. Correct me if I'm 
wrong, but it doesn't seem like the output needs to have any "ML Attribute" 
metadata.

Author: Nick Pritchard <[email protected]>

Closes #8751 from pnpritchard/SPARK-10573.

(cherry picked from commit 8a634e9bcc671167613fb575c6c0c054fb4b3479)
Signed-off-by: Xiangrui Meng <[email protected]>


Project: http://git-wip-us.apache.org/repos/asf/spark/repo
Commit: http://git-wip-us.apache.org/repos/asf/spark/commit/5b7067c9
Tree: http://git-wip-us.apache.org/repos/asf/spark/tree/5b7067c9
Diff: http://git-wip-us.apache.org/repos/asf/spark/diff/5b7067c9

Branch: refs/heads/branch-1.5
Commit: 5b7067c91f359356c5f65ea679d47dd6bf8b2eac
Parents: 5f58704
Author: Nick Pritchard <[email protected]>
Authored: Mon Sep 14 13:27:45 2015 -0700
Committer: Xiangrui Meng <[email protected]>
Committed: Mon Sep 14 13:31:58 2015 -0700

----------------------------------------------------------------------
 .../scala/org/apache/spark/ml/feature/StringIndexer.scala |  5 ++---
 .../org/apache/spark/ml/feature/StringIndexerSuite.scala  | 10 +++++++++-
 2 files changed, 11 insertions(+), 4 deletions(-)
----------------------------------------------------------------------


http://git-wip-us.apache.org/repos/asf/spark/blob/5b7067c9/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala
----------------------------------------------------------------------
diff --git 
a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala 
b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala
index 8a74da5..048b5e8 100644
--- a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala
+++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala
@@ -27,7 +27,7 @@ import org.apache.spark.ml.Transformer
 import org.apache.spark.ml.util.Identifiable
 import org.apache.spark.sql.DataFrame
 import org.apache.spark.sql.functions._
-import org.apache.spark.sql.types.{DoubleType, NumericType, StringType, 
StructType}
+import org.apache.spark.sql.types._
 import org.apache.spark.util.collection.OpenHashMap
 
 /**
@@ -220,8 +220,7 @@ class IndexToString private[ml] (
     val outputColName = $(outputCol)
     require(inputFields.forall(_.name != outputColName),
       s"Output column $outputColName already exists.")
-    val attr = NominalAttribute.defaultAttr.withName($(outputCol))
-    val outputFields = inputFields :+ attr.toStructField()
+    val outputFields = inputFields :+ StructField($(outputCol), StringType)
     StructType(outputFields)
   }
 

http://git-wip-us.apache.org/repos/asf/spark/blob/5b7067c9/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala
----------------------------------------------------------------------
diff --git 
a/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala 
b/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala
index 5fe66a3..ad008b0 100644
--- a/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala
+++ b/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala
@@ -17,7 +17,8 @@
 
 package org.apache.spark.ml.feature
 
-import org.apache.spark.SparkFunSuite
+import org.apache.spark.sql.types.{StringType, StructType, StructField, 
DoubleType}
+import org.apache.spark.{SparkException, SparkFunSuite}
 import org.apache.spark.ml.attribute.{Attribute, NominalAttribute}
 import org.apache.spark.ml.param.ParamsSuite
 import org.apache.spark.ml.util.MLTestingUtils
@@ -134,4 +135,11 @@ class StringIndexerSuite extends SparkFunSuite with 
MLlibTestSparkContext {
         assert(a === b)
     }
   }
+
+  test("IndexToString.transformSchema (SPARK-10573)") {
+    val idxToStr = new 
IndexToString().setInputCol("input").setOutputCol("output")
+    val inSchema = StructType(Seq(StructField("input", DoubleType)))
+    val outSchema = idxToStr.transformSchema(inSchema)
+    assert(outSchema("output").dataType === StringType)
+  }
 }


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to