Skip to content
Closed
Show file tree
Hide file tree
Changes from 5 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
160 changes: 156 additions & 4 deletions mllib/src/main/scala/org/apache/spark/ml/Pipeline.scala
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,19 @@ import java.{util => ju}
import scala.collection.JavaConverters._
import scala.collection.mutable.ListBuffer

import org.apache.spark.Logging
import org.apache.hadoop.fs.Path
import org.json4s._
import org.json4s.jackson.JsonMethods._

import org.apache.spark.{SparkContext, Logging}
import org.apache.spark.annotation.{DeveloperApi, Experimental}
import org.apache.spark.ml.param.{Param, ParamMap, Params}
import org.apache.spark.ml.util.Identifiable
import org.apache.spark.ml.util.Reader
import org.apache.spark.ml.util.Writer
import org.apache.spark.ml.util._
import org.apache.spark.sql.DataFrame
import org.apache.spark.sql.types.StructType
import org.apache.spark.util.Utils

/**
* :: DeveloperApi ::
Expand Down Expand Up @@ -82,7 +89,7 @@ abstract class PipelineStage extends Params with Logging {
* an identity transformer.
*/
@Experimental
class Pipeline(override val uid: String) extends Estimator[PipelineModel] {
class Pipeline(override val uid: String) extends Estimator[PipelineModel] with Writable {

def this() = this(Identifiable.randomUID("pipeline"))

Expand Down Expand Up @@ -166,6 +173,34 @@ class Pipeline(override val uid: String) extends Estimator[PipelineModel] {
"Cannot have duplicate components in a pipeline.")
theStages.foldLeft(schema)((cur, stage) => stage.transformSchema(cur))
}

override def write: Writer = new PipelineWriter(this)
}

object Pipeline extends Readable[Pipeline] {

override def read: Reader[Pipeline] = new PipelineReader

override def load(path: String): Pipeline = read.load(path)
}

private[ml] class PipelineWriter(instance: Pipeline) extends Writer {

PipelineSharedWriter.validateStages(instance.getStages)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should users be able to save an incomplete pipeline? For example, I could make a template pipeline, send it to other users, and they only need to fill in some required params like inputCol after they load it back.


override protected def saveImpl(path: String): Unit =
PipelineSharedWriter.saveImpl(instance, instance.getStages, sc, path)
}

private[ml] class PipelineReader extends Reader[Pipeline] {

/** Checked against metadata when loading model */
private val className = "org.apache.spark.ml.Pipeline"

override def load(path: String): Pipeline = {
val (uid: String, stages: Array[PipelineStage]) = PipelineSharedReader.load(className, sc, path)
new Pipeline(uid).setStages(stages)
}
}

/**
Expand All @@ -176,7 +211,7 @@ class Pipeline(override val uid: String) extends Estimator[PipelineModel] {
class PipelineModel private[ml] (
override val uid: String,
val stages: Array[Transformer])
extends Model[PipelineModel] with Logging {
extends Model[PipelineModel] with Writable with Logging {

/** A Java/Python-friendly auxiliary constructor. */
private[ml] def this(uid: String, stages: ju.List[Transformer]) = {
Expand All @@ -200,4 +235,121 @@ class PipelineModel private[ml] (
override def copy(extra: ParamMap): PipelineModel = {
new PipelineModel(uid, stages.map(_.copy(extra))).setParent(parent)
}

override def write: Writer = new PipelineModelWriter(this)
}

object PipelineModel extends Readable[PipelineModel] {

override def read: Reader[PipelineModel] = new PipelineModelReader

override def load(path: String): PipelineModel = read.load(path)
}

private[ml] class PipelineModelWriter(instance: PipelineModel) extends Writer {

PipelineSharedWriter.validateStages(instance.stages.asInstanceOf[Array[PipelineStage]])

override protected def saveImpl(path: String): Unit = PipelineSharedWriter.saveImpl(instance,
instance.stages.asInstanceOf[Array[PipelineStage]], sc, path)
}

private[ml] class PipelineModelReader extends Reader[PipelineModel] {

/** Checked against metadata when loading model */
private val className = "org.apache.spark.ml.PipelineModel"

override def load(path: String): PipelineModel = {
val (uid: String, stages: Array[PipelineStage]) =
PipelineSharedReader.load(className, sc, path)
val transformers = stages map {
case stage: Transformer => stage
case stage => throw new RuntimeException(s"PipelineModel.read loaded a stage but found it" +

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

minor: stage -> other

s" was not a Transformer. Bad stage: ${stage.uid}")

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

include stage.getClass in the error message

}
new PipelineModel(uid, transformers)
}
}

/** Methods for [[Writer]] shared between [[Pipeline]] and [[PipelineModel]] */
private[ml] object PipelineSharedWriter {

import org.json4s.JsonDSL._

/** Check that all stages are Writable */
def validateStages(stages: Array[PipelineStage]): Unit = {
stages.foreach {
case stage: Writable => // good
case stage =>

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

minor: other

throw new UnsupportedOperationException("Pipeline write will fail on this Pipeline" +

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

minor: will fail -> failed, this Pipeline -> pipeline ${pipeline.uid}

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

But a user could write:

val writer = pipeline.write   // failure will occur here, before attempting to write
writer.save(...)

s" because it contains a stage which does not implement Writable. Non-Writable stage:" +
s" ${stage.uid}")

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Include stage.getClass. Though the default implementation of uid contains class name, user can set arbitrary values.

}
}

def saveImpl(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

minor: saveImpl -> save?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I like saveImpl since it's analogous to the other saveImpl methods which call it

instance: Params,
stages: Array[PipelineStage],
sc: SparkContext,
path: String): Unit = {
// Copied and edited from DefaultParamsWriter.saveMetadata
// TODO: modify DefaultParamsWriter.saveMetadata to avoid duplication
val uid = instance.uid
val cls = instance.getClass.getName
val stageUids = stages.map(_.uid)
val jsonParams = List("stageUids" -> parse(compact(render(stageUids.toSeq))))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

  • Pass stageUids directly if you define stageUids as stages.map(_.uid).toSeq. I guess it should work.
  • We might consider renaming stageUids to stages. stageUids is more accurate but stages might be sufficient. For simple transformers, we could embed them into the metdata JSON.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is another way than this better?

val metadata = ("class" -> cls) ~
("timestamp" -> System.currentTimeMillis()) ~
("sparkVersion" -> sc.version) ~
("uid" -> uid) ~
("paramMap" -> jsonParams)
val metadataPath = new Path(path, "metadata").toString
val metadataJson = compact(render(metadata))
sc.parallelize(Seq(metadataJson), 1).saveAsTextFile(metadataPath)

// Save stages
val stagesDir = new Path(path, "stages").toString
stages.foreach {
case stage: Writable =>
val stagePath = new Path(stagesDir, stage.uid).toString

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It is useful if we prefix the stagePath with index, e.g., 000_Tokeninzer_123/, 001_LogisticRegression_456/. This helps manual inspection.

stage.write.save(stagePath)
}
}
}

/** Methods for [[Reader]] shared between [[Pipeline]] and [[PipelineModel]] */
private[ml] object PipelineSharedReader {

def load(className: String, sc: SparkContext, path: String): (String, Array[PipelineStage]) = {
val metadata = DefaultParamsReader.loadMetadata(path, sc, className)

implicit val format = DefaultFormats
val stagesDir = new Path(path, "stages").toString
val stageUids: Array[String] = metadata.params match {
case JObject(pairs) =>
if (pairs.length != 1) {
// Should not happen unless file is corrupted or we have a bug.
throw new RuntimeException(
s"Pipeline read expected 1 Param (stageUids), but found ${pairs.length}.")
}
pairs.head match {
case ("stageUids", jsonValue) =>
parse(compact(render(jsonValue))).extract[Seq[String]].toArray

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Would jsonValue.extract[Seq[String]].toArray work?

case (paramName, jsonValue) =>
// Should not happen unless file is corrupted or we have a bug.
throw new RuntimeException(s"Pipeline read encountered unexpected Param $paramName" +
s" in metadata: ${metadata.metadataStr}")
}
case _ =>
throw new IllegalArgumentException(
s"Cannot recognize JSON metadata: ${metadata.metadataStr}.")
}
val stages: Array[PipelineStage] = stageUids.map { stageUid =>
val stagePath = new Path(stagesDir, stageUid).toString
val stageMetadata = DefaultParamsReader.loadMetadata(stagePath, sc)
val cls = Utils.classForName(stageMetadata.className)
cls.getMethod("read").invoke(cls).asInstanceOf[Reader[PipelineStage]].load(stagePath)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

invoke(null)

}
(metadata.uid, stages)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -161,6 +161,8 @@ trait Readable[T] {

/**
* Reads an ML instance from the input path, a shortcut of `read.load(path)`.
*
* Note: Implementing classes should override this to be Java-friendly.
*/
@Since("1.6.0")
def load(path: String): T = read.load(path)
Expand All @@ -187,7 +189,7 @@ private[ml] object DefaultParamsWriter {
* - timestamp
* - sparkVersion
* - uid
* - paramMap
* - paramMap: These must be encodable using [[org.apache.spark.ml.param.Param.jsonEncode()]].
*/
def saveMetadata(instance: Params, path: String, sc: SparkContext): Unit = {
val uid = instance.uid
Expand Down
94 changes: 91 additions & 3 deletions mllib/src/test/scala/org/apache/spark/ml/PipelineSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -25,11 +25,13 @@ import org.scalatest.mock.MockitoSugar.mock

import org.apache.spark.SparkFunSuite
import org.apache.spark.ml.feature.HashingTF
import org.apache.spark.ml.param.ParamMap
import org.apache.spark.ml.util.MLTestingUtils
import org.apache.spark.ml.param.{IntParam, ParamMap}
import org.apache.spark.ml.util._
import org.apache.spark.mllib.util.MLlibTestSparkContext
import org.apache.spark.sql.DataFrame
import org.apache.spark.sql.types.StructType

class PipelineSuite extends SparkFunSuite {
class PipelineSuite extends SparkFunSuite with MLlibTestSparkContext with DefaultReadWriteTest {

abstract class MyModel extends Model[MyModel]

Expand Down Expand Up @@ -111,4 +113,90 @@ class PipelineSuite extends SparkFunSuite {
assert(pipelineModel1.uid === "pipeline1")
assert(pipelineModel1.stages === stages)
}

test("Pipeline read/write") {
val writableStage = new WritableStage("writableStage").setIntParam(56)
val pipeline = new Pipeline().setStages(Array(writableStage))

val pipeline2 = testDefaultReadWrite(pipeline, testParams = false)
assert(pipeline2.getStages.length === 1)
assert(pipeline2.getStages(0).isInstanceOf[WritableStage])
val writableStage2 = pipeline2.getStages(0).asInstanceOf[WritableStage]
assert(writableStage.getIntParam === writableStage2.getIntParam)
}

test("Pipeline read/write with non-Writable stage") {
val unWritableStage = new UnWritableStage("unwritableStage")
val unWritablePipeline = new Pipeline().setStages(Array(unWritableStage))
withClue("Pipeline.write should fail when Pipeline contains non-Writable stage") {
intercept[UnsupportedOperationException] {
unWritablePipeline.write
}
}
}

test("PipelineModel read/write") {
val writableStage = new WritableStage("writableStage").setIntParam(56)
val pipeline =
new PipelineModel("pipeline_89329327", Array(writableStage.asInstanceOf[Transformer]))

val pipeline2 = testDefaultReadWrite(pipeline, testParams = false)
assert(pipeline2.stages.length === 1)
assert(pipeline2.stages(0).isInstanceOf[WritableStage])
val writableStage2 = pipeline2.stages(0).asInstanceOf[WritableStage]
assert(writableStage.getIntParam === writableStage2.getIntParam)
}

test("PipelineModel read/write with non-Writable stage") {
val unWritableStage = new UnWritableStage("unwritableStage")
val unWritablePipeline =
new PipelineModel("pipeline_328957", Array(unWritableStage.asInstanceOf[Transformer]))
withClue("PipelineModel.write should fail when PipelineModel contains non-Writable stage") {
intercept[UnsupportedOperationException] {
unWritablePipeline.write
}
}
}
}


/** Used to test [[Pipeline]] with [[Writable]] stages */
class WritableStage(override val uid: String) extends Transformer with Writable {

final val intParam: IntParam = new IntParam(this, "intParam", "doc")

def getIntParam: Int = $(intParam)

def setIntParam(value: Int): this.type = set(intParam, value)

setDefault(intParam -> 0)

override def copy(extra: ParamMap): WritableStage = defaultCopy(extra)

override def write: Writer = new DefaultParamsWriter(this)

override def transform(dataset: DataFrame): DataFrame = dataset

override def transformSchema(schema: StructType): StructType = schema
}

object WritableStage extends Readable[WritableStage] {

override def read: Reader[WritableStage] = new DefaultParamsReader[WritableStage]

override def load(path: String): WritableStage = read.load(path)
}

/** Used to test [[Pipeline]] with non-[[Writable]] stages */
class UnWritableStage(override val uid: String) extends Transformer {

final val intParam: IntParam = new IntParam(this, "intParam", "doc")

setDefault(intParam -> 0)

override def copy(extra: ParamMap): UnWritableStage = defaultCopy(extra)

override def transform(dataset: DataFrame): DataFrame = dataset

override def transformSchema(schema: StructType): StructType = schema
}
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,13 @@ trait DefaultReadWriteTest extends TempDirectory { self: Suite =>
/**
* Checks "overwrite" option and params.
* @param instance ML instance to test saving/loading
* @param testParams If true, then test values of Params. Otherwise, just test overwrite option.
* @tparam T ML instance type
* @return Instance loaded from file
*/
def testDefaultReadWrite[T <: Params with Writable](instance: T): T = {
def testDefaultReadWrite[T <: Params with Writable](
instance: T,
testParams: Boolean = true): T = {
val uid = instance.uid
val path = new File(tempDir, uid).getPath

Expand All @@ -46,16 +49,18 @@ trait DefaultReadWriteTest extends TempDirectory { self: Suite =>
val newInstance = loader.load(path)

assert(newInstance.uid === instance.uid)
instance.params.foreach { p =>
if (instance.isDefined(p)) {
(instance.getOrDefault(p), newInstance.getOrDefault(p)) match {
case (Array(values), Array(newValues)) =>
assert(values === newValues, s"Values do not match on param ${p.name}.")
case (value, newValue) =>
assert(value === newValue, s"Values do not match on param ${p.name}.")
if (testParams) {
instance.params.foreach { p =>
if (instance.isDefined(p)) {
(instance.getOrDefault(p), newInstance.getOrDefault(p)) match {
case (Array(values), Array(newValues)) =>
assert(values === newValues, s"Values do not match on param ${p.name}.")
case (value, newValue) =>
assert(value === newValue, s"Values do not match on param ${p.name}.")
}
} else {
assert(!newInstance.isDefined(p), s"Param ${p.name} shouldn't be defined.")
}
} else {
assert(!newInstance.isDefined(p), s"Param ${p.name} shouldn't be defined.")
}
}

Expand Down