-
Notifications
You must be signed in to change notification settings - Fork 29.3k
[SPARK-9906] [ML] User guide for LogisticRegressionSummary #8197
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 2 commits
487b361
56cb35b
9831270
1ab3d9c
9825b14
1dfe7f6
4060c5b
83d229f
7bf922c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -801,6 +801,146 @@ jsc.stop(); | |
|
|
||
| </div> | ||
|
|
||
| ## Examples: Summaries for LogisticRegression. | ||
|
|
||
| Once [`LogisticRegression`](api/scala/index.html#org.apache.spark.ml.classification.LogisticRegression) | ||
| is run on data, it is useful to extract statistics such as the | ||
| loss per iteration which will provide an intuition on overfitting and metrics to understand | ||
| how well the model has performed on training and test data. | ||
|
|
||
| [`LogisticRegressionTrainingSummary`](api/scala/index.html#org.apache.spark.mllib.classification.LogisticRegressionTrainingSummary) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Since this links the scala api doc, move this under the scala codetab and add another one for the java api doc
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. |
||
| provides an interface to access such relevant information. i.e the objectiveHistory and metrics | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. "
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Actually, |
||
| to evaluate the performance on the training data directly with very less code to be rewritten by | ||
| the user. In the future, a method would be made available in the fitted | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. nit: I would just leave this part of the doc out in case we forget to remove it when we do settle on a good name for |
||
| [`LogisticRegressionModel`](api/scala/index.html#org.apache.spark.ml.classification.LogisticRegressionModel) to obtain | ||
| a [`LogisticRegressionSummary`](api/scala/index.html#org.apache.spark.mllib.classification.LogisticRegressionSummary) | ||
| of the test data as well. | ||
|
|
||
| This examples illustrates the use of `LogisticRegressionTrainingSummary` on some toy data. | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. examples -> example |
||
|
|
||
| <div class="codetabs"> | ||
| <div data-lang="scala"> | ||
| {% highlight scala %} | ||
| import org.apache.spark.{SparkConf, SparkContext} | ||
| import org.apache.spark.ml.classification.{LogisticRegression, BinaryLogisticRegressionSummary} | ||
| import org.apache.spark.mllib.regression.LabeledPoint | ||
| import org.apache.spark.mllib.linalg.Vectors | ||
| import org.apache.spark.sql.{Row, SQLContext} | ||
|
|
||
| val conf = new SparkConf().setAppName("LogisticRegressionSummary") | ||
| val sc = new SparkContext(conf) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. No need for
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I copied the example code shown previously where all these are set.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. We are not always consistent, but I've been trying to discourage those. Basically, I want users to be able to copy and paste the code into the spark shell and have it work.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I realized that and I removed it. |
||
| val sqlContext = new SQLContext(sc) | ||
| import sqlContext.implicits._ | ||
|
|
||
| // Use some random data for demonstration. | ||
| // Note that the RDD of LabeledPoints can be converted to a dataframe directly. | ||
| val data = sc.parallelize(Array( | ||
| LabeledPoint(0.0, Vectors.dense(0.2, 4.5, 1.6)), | ||
| LabeledPoint(1.0, Vectors.dense(3.1, 6.8, 3.6)), | ||
| LabeledPoint(0.0, Vectors.dense(2.4, 0.9, 1.9)), | ||
| LabeledPoint(1.0, Vectors.dense(9.1, 3.1, 3.6)), | ||
| LabeledPoint(0.0, Vectors.dense(2.5, 1.9, 9.1))) | ||
| ) | ||
| val logRegDataFrame = data.toDF() | ||
|
|
||
| // Run Logistic Regression on your toy data. | ||
| // Since LogisticRegression is an estimator, it returns an instance of LogisticRegressionModel | ||
| // which is a transformer. | ||
| val logReg = new LogisticRegression() | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Use builder notation |
||
| logReg.setMaxIter(5) | ||
| logReg.setRegParam(0.01) | ||
| val logRegModel = logReg.fit(logRegDataFrame) | ||
|
|
||
| // Extract the summary directly from the returned LogisticRegressionModel instance. | ||
| val trainingSummary = logRegModel.summary | ||
|
|
||
| // Obtain the loss per iteration. This should decrease upto a certain point and | ||
| // then increase or show negligible change after this. | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Training loss should always decrease; overfitting only affects the loss on test-set |
||
| val objectiveHistory = trainingSummary.objectiveHistory | ||
| objectiveHistory.foreach(loss => println(loss)) | ||
|
|
||
| // Obtain the metrics useful to judge performance on test data. | ||
| val binarySummary = trainingSummary.asInstanceOf[BinaryLogisticRegressionSummary] | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Add comment that that this cast is OK because we have a binary classification problem
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Ditto for java example |
||
|
|
||
| // Obtain the receiver-operating characteristic as a dataframe and areaUnderROC. | ||
| val roc = binarySummary.roc | ||
| roc.show() | ||
| roc.select("FPR").show() | ||
| println(binarySummary.areaUnderROC) | ||
|
|
||
| // Obtain the threshold with the highest fMeasure. | ||
| val fMeasure = binarySummary.fMeasureByThreshold | ||
| val fScoreRDD = fMeasure.map { case Row(thresh: Double, fscore: Double) => (thresh, fscore) } | ||
| val (highThresh, highFScore) = fScoreRDD.fold((0.0, 0.0))((threshFScore1, threshFScore2) => { | ||
| if (threshFScore1._2 > threshFScore2._2) threshFScore1 else threshFScore2 | ||
| }) | ||
|
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. @feynmanliang Is there are a shorthand to do this directly in Spark SQL? Once I understand that I can update the Java Example as well.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. You might be able to get it done by doing an aggregation (e.g.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I had thought the user can choose the best threshold and re-run LogisticRegression on this In the first run, it will be a random sampling of the dataset and in the second run it will be the entire dataset. Do you still think it will not be useful?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The closest I could come up with is.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Nice! I think what you described is useful, but is outside the scope of
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. OK. I wouldn't complain because it makes my job easier. |
||
| {% endhighlight %} | ||
| </div> | ||
|
|
||
| <div data-lang="java"> | ||
| {% highlight java %} | ||
| import com.google.common.collect.Lists; | ||
|
|
||
| import org.apache.spark.api.java.JavaRDD; | ||
| import org.apache.spark.ml.classification.BinaryLogisticRegressionSummary; | ||
| import org.apache.spark.ml.classification.LogisticRegression; | ||
| import org.apache.spark.ml.classification.LogisticRegressionModel; | ||
| import org.apache.spark.ml.classification.LogisticRegressionTrainingSummary; | ||
| import org.apache.spark.mllib.regression.LabeledPoint; | ||
| import org.apache.spark.mllib.linalg.Vectors; | ||
| import org.apache.spark.sql.DataFrame; | ||
| import org.apache.spark.sql.Row; | ||
|
|
||
| // Use some random data for demonstration. | ||
| // Note that the RDD of LabeledPoints can be converted to a dataframe directly. | ||
| JavaRDD<LabeledPoint> data = sc.parallelize(Lists.newArrayList( | ||
| new LabeledPoint(0.0, Vectors.dense(0.2, 4.5, 1.6)), | ||
| new LabeledPoint(1.0, Vectors.dense(3.1, 6.8, 3.6)), | ||
| new LabeledPoint(0.0, Vectors.dense(2.4, 0.9, 1.9)), | ||
| new LabeledPoint(1.0, Vectors.dense(9.1, 3.1, 3.6)), | ||
| new LabeledPoint(0.0, Vectors.dense(2.5, 1.9, 9.1))) | ||
| ); | ||
| DataFrame logRegDataFrame = sql.createDataFrame(data, LabeledPoint.class); | ||
|
|
||
| // Run Logistic Regression on your toy data. | ||
| // Since LogisticRegression is an estimator, it returns an instance of LogisticRegressionModel | ||
| // which is a transformer. | ||
| LogisticRegression logReg = new LogisticRegression(); | ||
| logReg.setMaxIter(5); | ||
| logReg.setRegParam(0.01); | ||
| LogisticRegressionModel logRegModel = logReg.fit(logRegDataFrame); | ||
|
|
||
| // Extract the summary directly from the returned LogisticRegressionModel instance. | ||
| LogisticRegressionTrainingSummary trainingSummary = logRegModel.summary(); | ||
|
|
||
| // Obtain the loss per iteration. This should decrease upto a certain point and | ||
| // then increase or show negligible change after this. | ||
| double[] objectiveHistory = trainingSummary.objectiveHistory(); | ||
| for (double lossPerIteration: objectiveHistory) { | ||
| System.out.println(lossPerIteration); | ||
| } | ||
|
|
||
| // Obtain the metrics useful to judge performance on test data. | ||
| BinaryLogisticRegressionSummary binarySummary = (BinaryLogisticRegressionSummary) trainingSummary; | ||
|
|
||
| // Obtain the receiver-operating characteristic as a dataframe and areaUnderROC. | ||
| DataFrame roc = binarySummary.roc(); | ||
| roc.show(); | ||
| roc.select("FPR").show(); | ||
| System.out.println(binarySummary.areaUnderROC()); | ||
|
|
||
| // Obtain the threshold with the highest fMeasure. | ||
| DataFrame fMeasure = binarySummary.fMeasureByThreshold(); | ||
|
|
||
|
|
||
| {% highlight %} | ||
| </div> | ||
|
|
||
| </div> | ||
|
|
||
|
|
||
|
|
||
|
|
||
| # Dependencies | ||
|
|
||
| Spark ML currently depends on MLlib and has the same dependencies. | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Examples -> Example
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Summaries -> Summary (There is only one summary)
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
No period at end