diff --git a/mllib/src/test/java/org/apache/spark/ml/classification/JavaDecisionTreeClassifierSuite.java b/mllib/src/test/java/org/apache/spark/ml/classification/JavaDecisionTreeClassifierSuite.java index 7de90500d6380..4c051ee981f38 100644 --- a/mllib/src/test/java/org/apache/spark/ml/classification/JavaDecisionTreeClassifierSuite.java +++ b/mllib/src/test/java/org/apache/spark/ml/classification/JavaDecisionTreeClassifierSuite.java @@ -17,6 +17,8 @@ package org.apache.spark.ml.classification; +import java.io.File; +import java.io.IOException; import java.util.HashMap; import java.util.Map; @@ -28,11 +30,12 @@ import org.apache.spark.ml.tree.impl.TreeTests; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; +import org.apache.spark.util.Utils; public class JavaDecisionTreeClassifierSuite extends SharedSparkSession { @Test - public void runDT() { + public void runDT() throws IOException { int nPoints = 20; double A = 2.0; double B = -1.5; @@ -62,18 +65,15 @@ public void runDT() { model.depth(); model.toDebugString(); - /* - // TODO: Add test once save/load are implemented. SPARK-6725 File tempDir = Utils.createTempDir(System.getProperty("java.io.tmpdir"), "spark"); - String path = tempDir.toURI().toString(); + String path = new File(tempDir, "model").getPath(); try { - model3.save(sc.sc(), path); + model.save(path); DecisionTreeClassificationModel sameModel = - DecisionTreeClassificationModel.load(sc.sc(), path); - TreeTests.checkEqual(model3, sameModel); + DecisionTreeClassificationModel.load(path); + TreeTests.checkEqual(model, sameModel); } finally { Utils.deleteRecursively(tempDir); } - */ } } diff --git a/mllib/src/test/java/org/apache/spark/ml/classification/JavaGBTClassifierSuite.java b/mllib/src/test/java/org/apache/spark/ml/classification/JavaGBTClassifierSuite.java index 03ed9c46db632..a28a6a9995e63 100644 --- a/mllib/src/test/java/org/apache/spark/ml/classification/JavaGBTClassifierSuite.java +++ b/mllib/src/test/java/org/apache/spark/ml/classification/JavaGBTClassifierSuite.java @@ -17,6 +17,8 @@ package org.apache.spark.ml.classification; +import java.io.File; +import java.io.IOException; import java.util.HashMap; import java.util.Map; @@ -28,11 +30,12 @@ import org.apache.spark.ml.tree.impl.TreeTests; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; +import org.apache.spark.util.Utils; public class JavaGBTClassifierSuite extends SharedSparkSession { @Test - public void runDT() { + public void runDT() throws IOException { int nPoints = 20; double A = 2.0; double B = -1.5; @@ -67,17 +70,14 @@ public void runDT() { model.trees(); model.treeWeights(); - /* - // TODO: Add test once save/load are implemented. SPARK-6725 File tempDir = Utils.createTempDir(System.getProperty("java.io.tmpdir"), "spark"); - String path = tempDir.toURI().toString(); + String path = new File(tempDir, "model").getPath(); try { - model3.save(sc.sc(), path); - GBTClassificationModel sameModel = GBTClassificationModel.load(sc.sc(), path); - TreeTests.checkEqual(model3, sameModel); + model.save(path); + GBTClassificationModel sameModel = GBTClassificationModel.load(path); + TreeTests.checkEqual(model, sameModel); } finally { Utils.deleteRecursively(tempDir); } - */ } } diff --git a/mllib/src/test/java/org/apache/spark/ml/classification/JavaRandomForestClassifierSuite.java b/mllib/src/test/java/org/apache/spark/ml/classification/JavaRandomForestClassifierSuite.java index 5e0c29d6b3261..25d8f579e7b29 100644 --- a/mllib/src/test/java/org/apache/spark/ml/classification/JavaRandomForestClassifierSuite.java +++ b/mllib/src/test/java/org/apache/spark/ml/classification/JavaRandomForestClassifierSuite.java @@ -17,6 +17,8 @@ package org.apache.spark.ml.classification; +import java.io.File; +import java.io.IOException; import java.util.HashMap; import java.util.Map; @@ -30,11 +32,12 @@ import org.apache.spark.ml.tree.impl.TreeTests; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; +import org.apache.spark.util.Utils; public class JavaRandomForestClassifierSuite extends SharedSparkSession { @Test - public void runDT() { + public void runDT() throws IOException { int nPoints = 20; double A = 2.0; double B = -1.5; @@ -86,18 +89,15 @@ public void runDT() { model.treeWeights(); Vector importances = model.featureImportances(); - /* - // TODO: Add test once save/load are implemented. SPARK-6725 File tempDir = Utils.createTempDir(System.getProperty("java.io.tmpdir"), "spark"); - String path = tempDir.toURI().toString(); + String path = new File(tempDir, "model").getPath(); try { - model3.save(sc.sc(), path); + model.save(path); RandomForestClassificationModel sameModel = - RandomForestClassificationModel.load(sc.sc(), path); - TreeTests.checkEqual(model3, sameModel); + RandomForestClassificationModel.load(path); + TreeTests.checkEqual(model, sameModel); } finally { Utils.deleteRecursively(tempDir); } - */ } } diff --git a/mllib/src/test/java/org/apache/spark/ml/regression/JavaDecisionTreeRegressorSuite.java b/mllib/src/test/java/org/apache/spark/ml/regression/JavaDecisionTreeRegressorSuite.java index 3eb589646773f..d26ecc89028f6 100644 --- a/mllib/src/test/java/org/apache/spark/ml/regression/JavaDecisionTreeRegressorSuite.java +++ b/mllib/src/test/java/org/apache/spark/ml/regression/JavaDecisionTreeRegressorSuite.java @@ -17,6 +17,8 @@ package org.apache.spark.ml.regression; +import java.io.File; +import java.io.IOException; import java.util.HashMap; import java.util.Map; @@ -29,12 +31,13 @@ import org.apache.spark.ml.tree.impl.TreeTests; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; +import org.apache.spark.util.Utils; public class JavaDecisionTreeRegressorSuite extends SharedSparkSession { @Test - public void runDT() { + public void runDT() throws IOException { int nPoints = 20; double A = 2.0; double B = -1.5; @@ -64,17 +67,14 @@ public void runDT() { model.depth(); model.toDebugString(); - /* - // TODO: Add test once save/load are implemented. SPARK-6725 File tempDir = Utils.createTempDir(System.getProperty("java.io.tmpdir"), "spark"); - String path = tempDir.toURI().toString(); + String path = new File(tempDir, "model").getPath(); try { - model2.save(sc.sc(), path); - DecisionTreeRegressionModel sameModel = DecisionTreeRegressionModel.load(sc.sc(), path); - TreeTests.checkEqual(model2, sameModel); + model.save(path); + DecisionTreeRegressionModel sameModel = DecisionTreeRegressionModel.load(path); + TreeTests.checkEqual(model, sameModel); } finally { Utils.deleteRecursively(tempDir); } - */ } } diff --git a/mllib/src/test/java/org/apache/spark/ml/regression/JavaGBTRegressorSuite.java b/mllib/src/test/java/org/apache/spark/ml/regression/JavaGBTRegressorSuite.java index 0e8bbd8ed6714..f60374fda3c3b 100644 --- a/mllib/src/test/java/org/apache/spark/ml/regression/JavaGBTRegressorSuite.java +++ b/mllib/src/test/java/org/apache/spark/ml/regression/JavaGBTRegressorSuite.java @@ -17,6 +17,8 @@ package org.apache.spark.ml.regression; +import java.io.File; +import java.io.IOException; import java.util.HashMap; import java.util.Map; @@ -29,12 +31,13 @@ import org.apache.spark.ml.tree.impl.TreeTests; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; +import org.apache.spark.util.Utils; public class JavaGBTRegressorSuite extends SharedSparkSession { @Test - public void runDT() { + public void runDT() throws IOException { int nPoints = 20; double A = 2.0; double B = -1.5; @@ -68,17 +71,14 @@ public void runDT() { model.trees(); model.treeWeights(); - /* - // TODO: Add test once save/load are implemented. SPARK-6725 File tempDir = Utils.createTempDir(System.getProperty("java.io.tmpdir"), "spark"); - String path = tempDir.toURI().toString(); + String path = new File(tempDir, "model").getPath(); try { - model2.save(sc.sc(), path); - GBTRegressionModel sameModel = GBTRegressionModel.load(sc.sc(), path); - TreeTests.checkEqual(model2, sameModel); + model.save(path); + GBTRegressionModel sameModel = GBTRegressionModel.load(path); + TreeTests.checkEqual(model, sameModel); } finally { Utils.deleteRecursively(tempDir); } - */ } } diff --git a/mllib/src/test/java/org/apache/spark/ml/regression/JavaRandomForestRegressorSuite.java b/mllib/src/test/java/org/apache/spark/ml/regression/JavaRandomForestRegressorSuite.java index d504ccb1d4bc8..96fabfa5193ae 100644 --- a/mllib/src/test/java/org/apache/spark/ml/regression/JavaRandomForestRegressorSuite.java +++ b/mllib/src/test/java/org/apache/spark/ml/regression/JavaRandomForestRegressorSuite.java @@ -17,6 +17,8 @@ package org.apache.spark.ml.regression; +import java.io.File; +import java.io.IOException; import java.util.HashMap; import java.util.Map; @@ -31,12 +33,13 @@ import org.apache.spark.ml.tree.impl.TreeTests; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; +import org.apache.spark.util.Utils; public class JavaRandomForestRegressorSuite extends SharedSparkSession { @Test - public void runDT() { + public void runDT() throws IOException { int nPoints = 20; double A = 2.0; double B = -1.5; @@ -88,17 +91,14 @@ public void runDT() { model.treeWeights(); Vector importances = model.featureImportances(); - /* - // TODO: Add test once save/load are implemented. SPARK-6725 File tempDir = Utils.createTempDir(System.getProperty("java.io.tmpdir"), "spark"); - String path = tempDir.toURI().toString(); + String path = new File(tempDir, "model").getPath(); try { - model2.save(sc.sc(), path); - RandomForestRegressionModel sameModel = RandomForestRegressionModel.load(sc.sc(), path); - TreeTests.checkEqual(model2, sameModel); + model.save(path); + RandomForestRegressionModel sameModel = RandomForestRegressionModel.load(path); + TreeTests.checkEqual(model, sameModel); } finally { Utils.deleteRecursively(tempDir); } - */ } }