A fitted Spark ML model is only useful if another job, often on another cluster and months later, can load it and get the same predictions. Spark ML persistence is what makes that possible. Any stage that implements the writable interface can be saved with model.write().save(path) and restored with the matching load(path). The call is one line. Everything around it is where production systems go wrong: what is written, which versions and languages can read it back, what happens to custom stages, and how a new model replaces an old one without a scoring job ever seeing a half-written directory.
This article opens a saved model directory and explains every part of it. It works through a train, save, validate and load cycle in PySpark, covers custom stages and tuning models, and sets out the compatibility rules, failure modes and options for serving outside Spark. Building the pipeline itself is covered in Spark ML Pipelines, in depth; this page starts where fit() returns.
What gets saved: params and learned state
Spark ML has two kinds of stage. An Estimator learns from data; a Transformer maps one DataFrame to another. A fitted model, such as LogisticRegressionModel or a whole PipelineModel, is a Transformer. Persisting it means writing two things: its params, the configuration values such as regParam and featuresCol, and its learned state, such as coefficients, the label list of a StringIndexer or the nodes of a tree ensemble.
Spark writes them in different forms. Params go into a small JSON document. Learned state goes into Parquet, written through the same DataFrame writer as any other dataset. A stage with no learned state, such as VectorAssembler, has only the JSON. Because the writer is ordinary Spark I/O, the path can be anything the cluster's Hadoop file system layer understands: HDFS, S3 through s3a://, ADLS, GCS or a local directory in tests.
Unsaved Estimators can be persisted too. Saving an unfitted Pipeline records the recipe without the result, which is useful for reproducing training later. Most of the time, though, you save the fitted PipelineModel, because scoring needs the learned state of every stage, not only the final model.
Anatomy of a saved model directory
The layout below is from the Spark 4.0 source and has the same shape in 3.x. Knowing it lets you inspect a model without starting Spark, and diagnose a failed load from the files alone.
metadata/ holds one line of JSON in a part file. Its fields are class (the fully qualified class name used to load it), timestamp, sparkVersion (the version that wrote it), uid, paramMap (params set explicitly) and defaultParamMap (defaults in force at save time, recorded since Spark 2.4). Keeping explicit and default params apart matters across upgrades: if a new release changes a default, the loaded model keeps the default it was trained with.
stages/ exists for pipelines. The pipeline metadata lists stageUids in order, and each stage gets a subdirectory named by its index and uid, with the index zero-padded so the names sort correctly. Each subdirectory is a complete saved stage with its own metadata and, if it learned anything, its own data/.
data/ is Parquet. For a logistic regression it holds the class count, feature count, intercepts and coefficient matrix; for tree models, one row per node. It is readable with spark.read.parquet, which is the fastest way to check what a model actually learned.
Worked example: train, save, validate, promote
A churn model shows the full cycle. Training fits the pipeline, writes it to a new path named by run time, and writes a manifest beside it. Spark does not record the training data, code version or metrics, so you must.
from datetime import datetime, timezone
import json, subprocess
from pyspark.ml import Pipeline
from pyspark.ml.feature import StringIndexer, VectorAssembler
from pyspark.ml.classification import LogisticRegression
from pyspark.ml.functions import vector_to_array
BASE = "s3a://ml-models/churn"
run = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H%M")
path = f"{BASE}/v={run}"
train = spark.read.parquet("s3a://features/churn/snapshot=2026-10-01")
pipe = Pipeline(stages=[
StringIndexer(inputCol="plan", outputCol="plan_idx", handleInvalid="keep"),
VectorAssembler(inputCols=["plan_idx", "tenure_m", "tickets_90d"], outputCol="features"),
LogisticRegression(labelCol="churned", regParam=0.01),
])
model = pipe.fit(train)
model.write().save(path) # no overwrite(): a new path every run
sample = train.sample(fraction=0.001, seed=7) # golden rows for the parity check
model.transform(sample) \n .select(*sample.columns, vector_to_array("probability")[1].alias("p_train")) \n .write.parquet(f"{path}/_golden")
manifest = {
"spark_version": spark.version,
"train_data": "s3a://features/churn/snapshot=2026-10-01",
"git_sha": subprocess.check_output(["git", "rev-parse", "HEAD"]).decode().strip(),
"auc_holdout": 0.871, # from your evaluation step
}
spark.createDataFrame([(json.dumps(manifest),)], ["value"]) \
.coalesce(1).write.text(f"{path}/_MANIFEST")Validation runs as a separate job, ideally on the Spark version the scoring jobs use. It loads the model from storage, not from memory, scores a fixed sample and compares the result with the training job's own predictions on the same rows. A model that cannot round-trip should never be promoted.
from pyspark.ml import PipelineModel
from pyspark.sql import functions as F
from pyspark.ml.functions import vector_to_array
loaded = PipelineModel.load(path)
golden = spark.read.parquet(f"{path}/_golden") # inputs + p_train, written at training time
scored = loaded.transform(golden.drop("p_train")) \
.select("id", vector_to_array("probability")[1].alias("p_load"))
diff = golden.join(scored, "id").select(F.max(F.abs(F.col("p_train") - F.col("p_load")))).first()[0]
assert diff < 1e-6, f"round-trip drift {diff}"Promotion is a small write: the model path goes into a pointer such as s3a://ml-models/churn/current.txt, or into a registry entry. Scoring jobs read the pointer at start-up, then call PipelineModel.load on what it names. Rollback is writing the previous path back. The same principle of immutable versioned artifacts is set out in the pipelines article; the point here is that the pointer is the only mutable object in the system.
Custom stages and the class-name rule
Built-in stages persist themselves. Custom stages are where most persistence bugs come from, and the cause is usually the same: the loader recreates the object from the class name recorded in metadata, so that name has to be importable wherever the model is loaded.
In PySpark, a transformer that has only params, and no learned state, gets persistence by mixing in DefaultParamsWritable and DefaultParamsReadable.
# churn_features/stages.py -- in a package installed on driver and executors
from pyspark import keyword_only
from pyspark.ml import Transformer
from pyspark.ml.param.shared import HasInputCol, HasOutputCol
from pyspark.ml.util import DefaultParamsReadable, DefaultParamsWritable
from pyspark.sql import functions as F
class LogTransform(Transformer, HasInputCol, HasOutputCol,
DefaultParamsReadable, DefaultParamsWritable):
@keyword_only
def __init__(self, inputCol=None, outputCol=None):
super().__init__()
self._set(**self._input_kwargs)
def _transform(self, df):
return df.withColumn(self.getOutputCol(), F.log1p(F.col(self.getInputCol())))Three rules follow from how the loader works.
- Never define a custom stage in a notebook cell or script body. PySpark records the class as module name plus class name. Defined in the main script, the name is
__main__.LogTransform, and loading fails in any other process. Put stages in a versioned package and install it wherever models are loaded. - A pipeline with any Python-only stage can only be loaded from Python. PySpark marks such a pipeline with
"language": "Python"in its metadata and writes it through a Python-side writer. A Scala or Java job callingPipelineModel.loadon it fails. If JVM services must load the model, write custom stages in Scala, or express the logic with built-in stages such asSQLTransformer. - Learned state needs its own writer. The default writer saves params only. A custom Estimator whose model holds fitted values, such as per-category target means, must save them itself: in Scala, implement
MLWritablewith anMLWriterthat writes adata/directory, and a matchingMLReader. A quick alternative is to store small fitted values as params, which keeps them in JSON.
Tuning models and sub-models
Tuning produces models that wrap other models. A CrossValidatorModel or TrainValidationSplitModel saves its estimator, evaluator, parameter grid, average metrics and best model. By default it does not save the model fitted for every grid point and fold, even if collectSubModels was on during fitting. To keep them, set the writer option persistSubModels:
cv_model.write().option("persistSubModels", "true").save(f"{BASE}/cv/v={run}")
cv_model.bestModel.write().save(path) # what scoring actually needsIn practice, save the best model as its own artifact, because that is what scoring loads, and save the validator model separately only if you need to audit the search. A grid of 24 settings with 5 folds is 120 fitted pipelines; persisting them all can take longer than the search itself. Fold design and metric reading are covered in Spark ML Cross-Validation, in depth.
Compatibility across versions, languages and Connect
The Spark ML guide states the compatibility contract precisely. Across minor and patch versions, a model saved by one version can be loaded by a later one, and behaves identically except for bug fixes; any breaking change is meant to be listed in the release notes, and an unlisted one is treated as a bug. Across major versions, loading and behaviour are best-effort only. The guide also says there is no guarantee of a stable file format; what is designed to be backwards compatible is the loading code.
| Saved by | Loaded by | Expectation | What to do |
|---|---|---|---|
| 3.5.x | 3.5.y | Loads, identical output | Nothing special |
| 3.4 | 3.5 | Loads; behaviour same apart from bug fixes | Run the parity check |
| 3.5 | 4.0 | Best-effort | Re-validate every model, or retrain on 4.0 |
| 4.0 | 3.5 | Not supported | Never load newer models on older Spark |
| Scala stages | Python, or the reverse | Same format, loads | Fine for built-in stages |
| Pipeline with a Python-only stage | Scala or Java | Fails | Port the stage, or load from Python |
Spark Connect adds one more rule. In Spark 4.0, when the Spark Connect ML module is on the classpath, the loader resolves the recorded class name through its allow-list, a guard against loading arbitrary classes from a crafted model directory. Built-in stages are on it. Test loading custom stages on your 4.0 deployment before relying on it; the client-server split itself is described in Spark Connect architecture.
The allow-list points at a wider fact: loading a model runs code chosen by the files. The class name in metadata decides what gets instantiated. Treat model storage like a code repository, with write access limited to the training pipeline, and do not load model directories from untrusted sources.
Serving a saved model
Persistence decides where a model can run. There are four realistic options.
- Batch scoring in Spark. Load, transform, write predictions keyed by entity. This is the natural fit and needs nothing beyond this article.
- Streaming in Spark. Most fitted stages are row-wise, so the same
PipelineModelcan transform a Structured Streaming DataFrame. Load it once at query start; a new model version needs a query restart. - Precomputed scores in a key-value store. For millisecond lookups, score in batch and serve the results. No model loading happens on the request path.
- Export. Some Scala models can be written in another format through the general writer, for example
lrModel.write.format("pmml").save(path)forLinearRegressionModel. Coverage is narrow and does not include whole pipelines. The more common route is to read the coefficients fromdata/, re-implement the final model in the serving language, and run the same golden-sample parity test against it.
A model registry such as MLflow can sit on top. Its Spark flavour stores this same native format with extra metadata, so everything above, including the custom-stage rules, still applies.
Failure modes
| Symptom | Cause | Fix |
|---|---|---|
| IOException: path already exists | Saving to an existing path without overwrite() | Write a new versioned path instead |
| Scoring job fails mid-run after a retrain | overwrite() deletes the old directory before writing the new one | Never overwrite a path a reader uses; promote by pointer |
| ModuleNotFoundError or class not found on load | Custom stage defined in __main__ or package missing on the loader | Ship stages as an installed package on driver and executors |
| JVM load of a PySpark pipeline fails | Pipeline marked as Python because one stage is Python-only | Port the stage to Scala or built-ins |
| Predictions shift after a Spark upgrade | Bug fix or behaviour change across versions | Golden-sample parity check in the upgrade plan |
| Unseen category crashes scoring | StringIndexer saved with handleInvalid='error' | Set keep before fitting; it is saved with the model |
| Save takes many minutes | Validator saved with all sub-models, or a huge tree ensemble | Save bestModel only; check data/ size |
| Partial directory after a crash | Job died during save | Promote only after validation passes, never on save success alone |
Two of these deserve emphasis. First, overwrite() is not atomic: the Spark source deletes the existing directory and then writes, with a comment noting that it does not restore the old content if the save fails. A crash between the two leaves no model at all. Second, handleInvalid and similar settings are params, so they are frozen into the saved model. Feature handling for nulls and unseen values is discussed in Spark ML Feature Engineering, in depth.
What to do next
- Open one of your saved models and read the part file in
metadata/and each stage's metadata. Confirm thesparkVersionand class names are what you expect. - Change your training job to write a new versioned path every run, with a manifest naming the data snapshot, code version, Spark version and holdout metrics.
- Remove every
overwrite()on a path that scoring jobs read, and introduce a pointer file or registry entry as the only mutable reference. - Add a validation job that loads the saved model from storage, scores a golden sample and fails on drift above a tolerance. Gate promotion on it.
- Move any custom stages into an installed package; grep your model metadata for
__main__to find the ones that will break. - Before the next Spark upgrade, run the parity check for every live model on the new version, and plan retraining for any that cross a major version.