Every piece of Spark's DataFrame-based machine learning library, from StringIndexer to GBTClassifier, is built from two abstractions. A Transformer turns one DataFrame into another. An Estimator learns something from a DataFrame and returns a Model, and that Model is itself a Transformer. Once you understand those two contracts and the Params system underneath them, built-in stages stop being black boxes, and you can write your own stages that tune, persist and score exactly like the built-ins.

The broader pipeline story, including leakage and a simple stateless custom stage, is covered in Spark ML Pipelines, in depth. This article goes one level down. It looks at what fit and transform actually promise, how parameter values are resolved, how Pipeline.fit walks its stages, and then builds a complete learned stage: an Estimator that computes per-column quantile bounds and a Model that clips values to them, which survives save and load without a custom writer. The examples use PySpark. The Scala section covers the hooks that exist only on the JVM side.

The two contracts: transform and fit

A Transformer has one job: transform(dataset) returns a new DataFrame, normally the input plus one or more appended columns. It must not mutate the input, and it should be lazy. Because DataFrames are immutable query plans, a well-behaved transform adds expressions to the plan and does not trigger a Spark job. VectorAssembler, SQLTransformer and every fitted Model are Transformers.

An Estimator has one job too: fit(dataset) reads the data, usually triggering one or more jobs, and returns a Model. Everything learned, such as vocabularies, means, tree ensembles or coefficients, lives in the Model and not in the Estimator. The Estimator remains a reusable recipe. You can fit it on January data and on February data and get two independent Models, and the Estimator object is unchanged.

That split is the point of the design. Learned state is captured once, from training data only, and replayed at scoring time. Compute a mean inside a Transformer instead and scoring recomputes it on the scoring batch, so the feature silently changes meaning in production. If a step needs statistics from data, it is an Estimator plus a Model; if it is a pure function of each row, it is a Transformer.

Params: the configuration system underneath

A Param is a descriptor with a parent, a name, a doc string and an optional type converter. The value is not stored on the Param. It is stored in maps on the stage instance. Each instance keeps two maps, one for defaults and one for values the user set explicitly, and lookups go through getOrDefault.

from pyspark.ml.feature import Imputer

imp = Imputer(inputCols=["age"], outputCols=["age_f"])
imp.isSet(imp.strategy)        # False: nobody set it
imp.isDefined(imp.strategy)    # True: there is a default
imp.getOrDefault(imp.strategy) # 'mean'
print(imp.explainParams())     # every param, its doc, default and current value

Resolution follows a fixed precedence, and most tuning surprises come from forgetting it:

  1. An extra param map passed at call time, as in est.fit(df, {est.maxIter: 50}), wins. It is applied to a copy of the stage, so the original stage is not modified.
  2. A value set explicitly on the instance, through the constructor or a setter, comes next.
  3. The default declared by the stage is used only if neither of the above exists.

Two consequences follow. copy(extra) is how tuners create per-trial variants, so a custom stage must keep all configuration in Params; a plain attribute such as self.threshold = 3 is not copied, not explained and not saved. And type converters such as TypeConverters.toFloat run when a value is set, so a wrong type fails at configuration time rather than mid-job.

How Pipeline.fit walks its stages

A Pipeline is itself an Estimator whose stages are Transformers or Estimators. Its fit method is short and worth knowing exactly. It finds the index of the last Estimator in the stage list. Then it walks the stages in order. A Transformer is kept as is and applied to the running DataFrame. An Estimator is fitted on the running DataFrame, and the resulting Model is applied only if there are more Estimators after it. The output is a PipelineModel holding the original Transformers and the fitted Models in order.

Pipeline.fit: estimators learn on the data as transformed by every earlier stagetrain DataFrameraw columnsStage 1TransformerStage 2EstimatorStage 3Estimator (last)transformdf1df2Model 2learned stateModel 3learned statefit(df1)fit(df2)Model 2 transforms df1 into df2 because another estimator follows it.Model 3 is never applied during fit.PipelineModel = [Stage 1, Model 2, Model 3]a Transformer: one transform() call replays every stage at scoring timeFit-time and score-time code paths are the same stage objects, which is what prevents training-serving skew.
Each Estimator sees the output of all earlier stages. The final PipelineModel replays the same stages for scoring.

Three points follow. A scaler placed after an imputer learns statistics of imputed values. Stages after the last Estimator do not run during fit, so an error in a trailing Transformer only appears at transform. And each fit is a separate action, so five Estimators read the input five or more times: cache an expensive training DataFrame before fitting and unpersist it afterwards.

Worked example: a quantile-clipping Estimator and Model

Here is a learned stage that comes up in many tabular projects: winsorising, meaning clipping each numeric column to its own lower and upper quantile so that a few extreme values do not dominate a linear model. The bounds must come from training data, so this is an Estimator that produces a Model. The design goals are to support many columns in one pass, to keep all learned state in Params so default persistence works, and to use only built-in column expressions so the work stays in the JVM.

from pyspark import keyword_only
from pyspark.ml import Estimator, Model
from pyspark.ml.param import Param, Params, TypeConverters
from pyspark.ml.param.shared import HasInputCols, HasOutputCols
from pyspark.ml.util import DefaultParamsReadable, DefaultParamsWritable
from pyspark.sql import functions as F


class _ClipperParams(HasInputCols, HasOutputCols):
    lowerQuantile = Param(Params._dummy(), "lowerQuantile",
                          "quantile used as the lower bound, in [0, 1)",
                          typeConverter=TypeConverters.toFloat)
    upperQuantile = Param(Params._dummy(), "upperQuantile",
                          "quantile used as the upper bound, in (0, 1]",
                          typeConverter=TypeConverters.toFloat)
    relativeError = Param(Params._dummy(), "relativeError",
                          "relative error passed to approxQuantile",
                          typeConverter=TypeConverters.toFloat)

    def __init__(self, *args):
        super().__init__(*args)
        self._setDefault(lowerQuantile=0.01, upperQuantile=0.99,
                         relativeError=0.001)


class QuantileClipper(Estimator, _ClipperParams,
                      DefaultParamsReadable, DefaultParamsWritable):
    @keyword_only
    def __init__(self, inputCols=None, outputCols=None, lowerQuantile=0.01,
                 upperQuantile=0.99, relativeError=0.001):
        super().__init__()
        self._set(**self._input_kwargs)

    def _fit(self, dataset):
        cols = self.getInputCols()
        lo_q = self.getOrDefault(self.lowerQuantile)
        hi_q = self.getOrDefault(self.upperQuantile)
        if not 0.0 <= lo_q < hi_q <= 1.0:
            raise ValueError(f"need 0 <= lower < upper <= 1, got {lo_q}, {hi_q}")
        if len(cols) != len(self.getOutputCols()):
            raise ValueError("inputCols and outputCols differ in length")
        # One job for all columns; nulls are ignored by approxQuantile.
        q = dataset.approxQuantile(cols, [lo_q, hi_q],
                                   self.getOrDefault(self.relativeError))
        empty = [c for c, b in zip(cols, q) if len(b) < 2]
        if empty:
            raise ValueError(f"no non-null values to learn bounds from: {empty}")
        model = self._copyValues(QuantileClipperModel())
        return model._set(lowerBounds=[b[0] for b in q],
                          upperBounds=[b[1] for b in q])


class QuantileClipperModel(Model, _ClipperParams,
                           DefaultParamsReadable, DefaultParamsWritable):
    lowerBounds = Param(Params._dummy(), "lowerBounds",
                        "learned lower bound per input column",
                        typeConverter=TypeConverters.toListFloat)
    upperBounds = Param(Params._dummy(), "upperBounds",
                        "learned upper bound per input column",
                        typeConverter=TypeConverters.toListFloat)

    def __init__(self):
        # No required arguments: the default reader calls cls() and then
        # restores every param from the saved metadata.
        super().__init__()

    def _transform(self, dataset):
        los = self.getOrDefault(self.lowerBounds)
        his = self.getOrDefault(self.upperBounds)
        return dataset.withColumns({
            out: F.least(F.greatest(F.col(inp), F.lit(lo)), F.lit(hi))
            for inp, out, lo, hi in zip(self.getInputCols(),
                                        self.getOutputCols(), los, his)
        })

The shared _ClipperParams mixin declares Params once at class level, so both classes expose the same names and _copyValues carries the user's settings to the Model. _fit validates configuration before spending cluster time, computes every quantile in one approxQuantile job, and rejects an all-null column instead of indexing into an empty result. The Model's _transform is one withColumns projection, so Catalyst folds it into the plan and no Python worker sees a row.

Using it looks exactly like using a built-in stage:

clip = QuantileClipper(inputCols=["spend", "sessions"],
                       outputCols=["spend_c", "sessions_c"],
                       upperQuantile=0.995)
model = clip.fit(train)
model.getOrDefault(model.upperBounds)   # e.g. [1843.2, 77.0]
scored = model.transform(test)          # lazy: no job until an action

Persisting learned state

The default writer saves a stage as a small metadata document containing its class name, uid and every set and default Param, encoded as JSON. The default reader imports the class by name, calls it with no arguments, restores the uid and sets the Params back. That explains the two rules in the Model above: the constructor must accept no required arguments, and every piece of learned state must be a Param with a JSON-friendly type.

State in Params is right when it is small: a few numbers per column, a short category list, a threshold. Large state, such as a million-term vocabulary, bloats the metadata and the driver memory needed to parse it, and non-JSON state such as a matrix may not round-trip. For those, save the state as a DataFrame under the model directory with a custom writer and reader, as built-in Models do. Those hooks vary by version, so copy the pattern from a built-in model in your installed PySpark.

Two more persistence facts matter for deployment. Any pipeline that includes a stage written in Python loads only from Python, with the defining module importable under the same name, a point covered in Spark ML Model Persistence, in depth. And Spark 4.0 made ML over Spark Connect generally available for Python, but the published material covers built-in algorithms; test a custom stage under Connect rather than assume it works.

Tuning a custom stage

Because the clipper exposes real Params, a tuner can search over them like any other hyperparameter. ParamGridBuilder produces a list of param maps, and CrossValidator calls fitMultiple on the pipeline, which in turn copies each stage with the trial's overrides.

# Pipeline, VectorAssembler, LinearRegression, CrossValidator, ParamGridBuilder
# and RegressionEvaluator come from pyspark.ml and its submodules.
clip = QuantileClipper(inputCols=["spend", "sessions"],
                       outputCols=["spend_c", "sessions_c"])
vec = VectorAssembler(inputCols=["spend_c", "sessions_c"], outputCol="features")
lr = LinearRegression(labelCol="ltv")
pipe = Pipeline(stages=[clip, vec, lr])

grid = (ParamGridBuilder()
        .addGrid(clip.upperQuantile, [0.95, 0.99, 0.999])
        .addGrid(lr.regParam, [0.0, 0.1])
        .build())
cv = CrossValidator(estimator=pipe, estimatorParamMaps=grid,
                    evaluator=RegressionEvaluator(labelCol="ltv"),
                    numFolds=3, parallelism=2)
best = cv.fit(train).bestModel   # a PipelineModel

Because the clipper sits inside the pipeline, each fold learns its bounds from that fold's training split only. That is the leakage protection you lose if you clip the whole dataset first and tune afterwards. The grid above has six combinations and three folds, so it runs eighteen pipeline fits and each one triggers a quantile job. Budget for it, and see Spark ML Cross-Validation, in depth for how parallelism and caching interact.

The Scala side: transformSchema and copy

On the JVM, a stage has one more hook that Python stages do not: transformSchema. Built-in stages implement it to check input column types and compute the output schema without touching data, and they call it at the start of fit and transform. That is why a wrong input type on a built-in fails fast with a clear message, while a Python stage only fails when the job runs. In Python, do the equivalent check yourself at the top of _fit and _transform by inspecting dataset.schema.

class QuantileClipper(override val uid: String)
    extends Estimator[QuantileClipperModel] with ClipperParams
    with DefaultParamsWritable {

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

  override def transformSchema(schema: StructType): StructType = {
    $(inputCols).foreach { c =>
      require(schema(c).dataType.isInstanceOf[NumericType], s"$c is not numeric")
    }
    $(outputCols).foldLeft(schema)((s, c) => s.add(c, DoubleType, nullable = true))
  }

  override def fit(dataset: Dataset[_]): QuantileClipperModel = {
    transformSchema(dataset.schema, logging = true)
    val q = dataset.stat.approxQuantile($(inputCols),
      Array($(lowerQuantile), $(upperQuantile)), $(relativeError))
    copyValues(new QuantileClipperModel(uid, q.map(_(0)), q.map(_(1))).setParent(this))
  }

  override def copy(extra: ParamMap): QuantileClipper = defaultCopy(extra)
}

Here the Model takes its learned arrays as constructor arguments, as the built-ins do, so it needs its own writer and reader. And copy is mandatory, usually via defaultCopy, which needs a constructor taking only the uid.

Failure modes

SymptomCauseFix
Loaded model has default or empty boundsLearned state kept in plain attributes, not ParamsMove state into Params or write a custom writer and reader
Load fails with a constructor errorModel __init__ has required argumentsMake every constructor argument optional
Tuning grid has no effect on a custom stageValue read from an attribute instead of getOrDefaultAlways read through getOrDefault so copies with overrides apply
Fold scores look too goodStatistics computed on the full dataset before CrossValidatorPut every learned step inside the pipeline being tuned
Scoring is much slower than trainingPython UDF inside _transformUse column expressions, or a pandas UDF if Python is unavoidable
Pipeline cannot load in a Scala serving jobIt contains Python-defined stagesImplement the stage in Scala, or serve from Python

A further quiet failure is non-determinism. If a stage samples, shuffles or uses random initialisation, expose a seed Param through the shared HasSeed mixin so that refits are reproducible.

Trade-offs and design rules

A custom stage is code you maintain and keep importable wherever the model loads. First check whether SQLTransformer or a combination of built-ins expresses the step, as discussed in Spark ML Feature Engineering, in depth. Built-ins load from any language, get schema validation and are maintained by the Spark project.

When a custom stage is justified, make it do one thing, take lists of columns so a wide table needs one job, keep learned state in Params when it is small, and prefer expressions to UDFs; if Python is unavoidable, read Pandas UDFs in Spark, in depth first.

What to do next

  1. List the learned steps in your current feature code, such as means, quantiles or vocabularies, and confirm each one is an Estimator inside the pipeline, not a statistic computed beforehand.
  2. Print explainParams() for every stage in one production pipeline and check that no setting you rely on is coming from an unnoticed default.
  3. For each custom stage, move any plain attributes into Params and make the Model constructor argument-free.
  4. Add a unit test that fits on a small local DataFrame, saves to a temporary directory, loads, transforms and asserts the outputs are equal.
  5. Replace any Python UDF in a _transform with column expressions where possible, and measure scoring time before and after.
  6. If you plan to use Spark Connect, run that same round-trip test through a Connect session before relying on it.
Key takeaway: A Transformer maps a DataFrame to a DataFrame; an Estimator learns from one and returns a Model, which is a Transformer. Keep all configuration and small learned state in Params, read them through getOrDefault, make Model constructors argument-free, put every learned step inside the pipeline you tune, and test the fit, save, load and transform round trip.