Spark's machine learning library, MLlib, lets you train models on data that is already in Spark, using the same cluster, the same DataFrames and the same jobs you use for ETL. That is its entire value proposition, and it is also its limit. Spark ML is excellent at classical, data-parallel algorithms (linear models, tree ensembles, clustering, matrix factorisation) over large tabular datasets. It is not a deep learning framework, and for data that fits on one machine it is usually slower than scikit-learn.

This article is the map. It explains what the library contains, how training actually executes on a cluster (which is what determines speed and failure), walks through a complete model in PySpark, and shows when to reach for a different tool. The pipeline contract (Transformers, Estimators, leak-free tuning, custom stages) has its own article, Spark ML Pipelines in depth; read it after this one.

Advertisement

spark.ml and spark.mllib: two APIs, one library

There are two packages. spark.ml (pyspark.ml in Python) is the DataFrame-based API and the one to use. spark.mllib is the original RDD-based API. Since Spark 2.0 the RDD-based API has been in maintenance mode: it receives bug fixes but no new features. It is not deprecated, and some algorithms and utilities exist only there, but new code should start from spark.ml.

The name MLlib covers both, which is why documentation and blog posts use it loosely. When someone says Spark MLlib today they almost always mean the DataFrame API. The practical differences: spark.ml works on DataFrames with a single vector column of features, composes stages into Pipelines, benefits from Catalyst for the feature preparation, and has consistent Python, Scala, Java and R APIs.

What is in the box

FamilyMain estimatorsTypical use
Feature engineeringVectorAssembler, StringIndexer, OneHotEncoder, StandardScaler, Bucketizer, Imputer, HashingTF, IDF, Word2VecTurn columns into the single features vector every estimator expects
ClassificationLogisticRegression, DecisionTree, RandomForest, GBTClassifier, LinearSVC, NaiveBayes, MultilayerPerceptronChurn, fraud, propensity scoring
RegressionLinearRegression, GeneralizedLinearRegression, RandomForest, GBTRegressor, AFTSurvivalRegression, IsotonicRegressionDemand, time-to-event, pricing
ClusteringKMeans, BisectingKMeans, GaussianMixture, LDA, PowerIterationClusteringSegmentation, topic modelling
Collaborative filteringALSRecommendations from implicit or explicit feedback
Pattern miningFPGrowth, PrefixSpanMarket-basket rules, sequence patterns
Tuning and evaluationCrossValidator, TrainValidationSplit, ParamGridBuilder, evaluatorsModel selection
StatisticsCorrelation, ChiSquareTest, SummarizerExploration and feature screening

Notice what is absent: no convolutional or transformer networks (the multilayer perceptron is a small fully connected classifier), no built-in gradient boosting as strong as XGBoost or LightGBM, and no online learning in the DataFrame API. Those gaps define when to use something else.

Advertisement

How training executes on a cluster

Understanding the execution model explains nearly every performance and failure question. Most Spark ML estimators are iterative: they make many passes over the training data. Each pass is an ordinary Spark job. The driver holds the model, the executors hold the data, and the two exchange small amounts of information per iteration.

Logistic regression uses quasi-Newton optimisers (L-BFGS, or OWL-QN when L1 regularisation is on). LinearRegression has a solver parameter: auto (the default) uses a one-pass normal-equation solver when the number of features is small and falls back to L-BFGS otherwise. In the iterative case, each iteration the driver broadcasts the current coefficients, every task computes the loss and gradient over its partition, and the partial results are combined with treeAggregate, which merges them in levels (aggregationDepth) rather than sending every partition's vector straight to the driver. The driver then takes an optimiser step. The cost per iteration is one scan of the data plus a gradient-sized network exchange, so wide feature vectors (millions of one-hot columns) make the aggregation, not the scan, the bottleneck.

Decision trees, random forests and gradient-boosted trees first discretise continuous features into at most maxBins bins. Trees are then grown level by level: for each node being split, executors build histograms of label statistics per feature and bin, the histograms are aggregated, and the driver chooses the best splits. GBTs train trees sequentially, so each boosting round is a full set of jobs; this is why a 500-round GBT is slow on Spark and why dedicated boosting libraries are popular.

ALS partitions users and items into blocks, alternately solving for user factors with item factors fixed and vice versa, shuffling factor blocks each half-iteration. Its lineage grows every iteration, which is why checkpointInterval exists: without checkpointing, long runs can hit stack overflows or very slow recovery.

One iteration of distributed training in Spark ML (linear models)Driverholds coefficients w_tbroadcast w_tExecutor 1partitions 1..k: sum gradExecutor 2partitions k+1..2kExecutor ncached training rowstreeAggregatepartial sums combined in levelsDriver: L-BFGS / OWL-QN stepw_t+1, check convergenceEach iteration is a Spark job. Cache the input, or every iteration re-reads and re-parses the source.
Data-parallel training of a linear model: the driver broadcasts the coefficients, executors compute partial gradients on cached partitions, treeAggregate combines them, and the driver takes an L-BFGS step.

Data representation: one vector column

Every estimator reads a single column, features by default, of type Vector, plus a label column of doubles for supervised learning. Vectors can be dense or sparse. VectorAssembler concatenates numeric columns and other vectors into that column; OneHotEncoder produces sparse vectors. Sparse representation matters: a one-hot encoding of 50,000 product ids stored densely would be 400 KB of doubles per row.

The linear algebra runs through Breeze and the dev.ludovic.netlib BLAS wrapper. Native BLAS libraries such as OpenBLAS or Intel MKL are not bundled with Spark; without them you get a warning about failing to load JNIBLAS and a pure JVM implementation. For dense, compute-heavy algorithms (ALS, large Gaussian mixtures) installing native BLAS on executors is worth measuring.

Vectors are a user-defined type, so ordinary SQL functions cannot see inside them. When you need to inspect or export features, convert with pyspark.ml.functions.vector_to_array and back with array_to_vector. Keep a record of the input column order passed to VectorAssembler: a coefficient or feature importance at index 17 is meaningless unless you can map it back to the column that produced it, and the ml_attr metadata on the output column is the most reliable place to recover that mapping.

Worked example: a churn model end to end

Suppose a subscription business has 40 million customer rows in Parquet with a plan type, tenure, monthly spend, ticket count and a churned flag. The job below prepares features, trains a regularised logistic regression, evaluates it on a held-out split and saves the fitted pipeline.

from pyspark.sql import SparkSession, functions as F
from pyspark.ml import Pipeline
from pyspark.ml.feature import StringIndexer, OneHotEncoder, VectorAssembler, StandardScaler
from pyspark.ml.classification import LogisticRegression
from pyspark.ml.evaluation import BinaryClassificationEvaluator

spark = SparkSession.builder.appName("churn").getOrCreate()
df = (spark.read.parquet("s3://lake/churn/features/")
      .withColumn("label", F.col("churned").cast("double"))
      .na.fill({"tenure_months": 0, "support_tickets": 0, "monthly_spend": 0.0}))

train, test = df.randomSplit([0.8, 0.2], seed=42)
train = train.repartition(200).cache()        # iterative algorithm: cache the input
train.count()                                 # materialise the cache once

pipe = Pipeline(stages=[
    StringIndexer(inputCol="plan", outputCol="plan_idx", handleInvalid="keep"),
    OneHotEncoder(inputCols=["plan_idx"], outputCols=["plan_vec"]),
    VectorAssembler(inputCols=["plan_vec", "tenure_months", "monthly_spend",
                               "support_tickets"], outputCol="raw",
                    handleInvalid="keep"),
    StandardScaler(inputCol="raw", outputCol="features"),
    LogisticRegression(maxIter=50, regParam=0.01, elasticNetParam=0.0,
                       aggregationDepth=2),
])
model = pipe.fit(train)

pred = model.transform(test)
auc = BinaryClassificationEvaluator(metricName="areaUnderROC").evaluate(pred)
print(f"test AUC = {auc:.3f}")
lr = model.stages[-1]
print("iterations:", lr.summary.totalIterations)
model.write().overwrite().save("s3://models/churn/v7")

Several choices here are deliberate. The training set is repartitioned and cached before fit because logistic regression will scan it once per iteration; without the cache, each of up to 50 iterations re-reads Parquet. handleInvalid='keep' on the indexer maps an unseen plan type at scoring time to an extra bucket instead of failing the job. On the assembler, keep turns a null into NaN inside the vector rather than an error, which trades a loud failure for NaN predictions, so nulls are filled explicitly first and the share of NaN predictions should be monitored. The whole pipeline, including the indexer's learned vocabulary and the scaler's statistics, is saved together, so scoring applies identical transformations. After training, check summary.totalIterations: hitting maxIter exactly means the optimiser did not converge and the coefficients may be poor.

A reasonable expectation for this shape of data is that training time is dominated by the number of iterations times the time to scan the cached data. If each scan takes 8 seconds and the model needs 40 iterations, you are near five and a half minutes regardless of how fast the optimiser step is on the driver.

Performance guidance

  • Cache the prepared training data, in memory or memory-and-disk, and materialise it once before fitting. Unpersist it afterwards. See Spark memory management for how storage and execution memory compete.
  • Right-size partitions. Too few leaves cores idle; too many adds per-task overhead to every iteration. A few partitions per core is a sane start.
  • Watch feature width. High-cardinality one-hot encodings make gradients and tree histograms large. Hash, bucket or drop rare categories before encoding.
  • Tune maxBins and maxDepth together. Histogram size grows with bins, features and the number of nodes at a level, which doubles each level.
  • Set checkpointing (spark.sparkContext.setCheckpointDir plus checkpointInterval) for ALS and long GBT runs.
  • Size the driver. Models, histogram aggregation and collected summaries live there; large-feature models fail with driver out-of-memory before executors are stressed.
  • Parallelise tuning. CrossValidator(parallelism=4) fits several parameter settings at once; the product of the grid and the folds is your real job count.

When to use something else

Use Spark ML when data is large, already in Spark, and the model is classical. Otherwise, pick the tool that fits and use Spark for what it does best: preparing data and fanning out work.

SituationBetter choiceHow Spark helps
Training data fits in one machine's memoryscikit-learn, LightGBM or XGBoost on one nodeSample or aggregate in Spark, then train locally
Gradient boosting at scaleXGBoost's xgboost.spark estimatorsDistributes training across executors with a Spark-ML-style API
Deep learningPyTorch via pyspark.ml.torch.distributor.TorchDistributor (Spark 3.4+)Launches distributed PyTorch on the cluster
Thousands of small per-group modelsgroupBy().applyInPandas with a single-node libraryParallelises independent fits
GPU-accelerated classical MLRAPIDS-based pluginsSee Spark on GPUs

The per-group pattern deserves an example because it is so often done wrongly, for example by looping over groups on the driver and calling fit thousands of times:

# Many small models (one per store) are better trained in parallel single-node jobs.
import pandas as pd
from sklearn.linear_model import Ridge

def fit_store(pdf: pd.DataFrame) -> pd.DataFrame:
    m = Ridge(alpha=1.0).fit(pdf[["price", "promo", "dow"]], pdf["units"])
    return pd.DataFrame({"store_id": [pdf.store_id.iloc[0]],
                         "coef_price": [m.coef_[0]], "r2": [m.score(
                             pdf[["price", "promo", "dow"]], pdf["units"])]})

models = (sales.groupBy("store_id")
          .applyInPandas(fit_store, "store_id long, coef_price double, r2 double"))

Spark ML over Spark Connect

Spark Connect separates a thin client from the cluster (the Spark Connect article explains the protocol). Spark 3.5 introduced a separate pyspark.ml.connect module with a small subset of algorithms. Spark 4.0 went further and allows the regular pyspark.ml estimators, transformers and pipelines to run over a Connect session, with the model objects living server-side and referenced from the client. If you are on a Connect-based platform, confirm the specific estimators you need against your platform's version before migrating; custom Python stages that assume a local SparkContext are the usual casualty.

Failure modes

  • Silent non-convergence. The optimiser stops at maxIter. Always log iterations and the objective history from the training summary.
  • Driver out of memory on wide models, large tree histograms or collecting predictions. Reduce feature width and never collect() scored data.
  • Recomputed lineage every iteration. Training is ten times slower than expected because the input was not cached, or was evicted.
  • Stack overflow or endless recovery in ALS. Missing checkpoint directory.
  • Train-serve skew. Features computed differently at scoring time. Persist and load the whole PipelineModel.
  • Label leakage in tuning. Fitting preprocessing on the full dataset before cross-validation; keep it inside the pipeline.
  • Version drift. Models saved by one Spark version are usually loadable by later ones, but test the load in CI before an upgrade.

What to do next

  1. Inventory your ML workloads and classify each as large-and-classical (Spark ML), small (single node), boosting (xgboost.spark), deep learning (TorchDistributor) or per-group (applyInPandas).
  2. Run the churn example on your own data, and record iterations, per-iteration time and whether the cache held.
  3. Move any preprocessing that happens outside the pipeline into pipeline stages and save the fitted PipelineModel.
  4. Check executor logs for the JNIBLAS warning and decide, by measurement, whether native BLAS is worth installing.
  5. Add checkpointing to every ALS and long GBT job.
  6. Read the pipelines deep dive next and adopt its leak-free tuning pattern.
Key takeaway: Spark ML is the DataFrame-based spark.ml package; the RDD-based spark.mllib is in maintenance mode but not deprecated. Its algorithms run as sequences of ordinary Spark jobs: linear models broadcast coefficients and treeAggregate gradients into an L-BFGS step, trees aggregate binned histograms level by level, and ALS alternates shuffled factor blocks. Performance follows from that model: cache the training input, control feature width, checkpoint long lineages and size the driver. Use it for large, classical, tabular problems already in Spark, and hand small data, boosting and deep learning to tools built for them, with Spark doing the preparation and fan-out.