MLlib is Apache Spark's machine-learning library. Most introductions focus on its classifiers and regressors, but in real projects most of the code, and most of the bugs, sit around the model: turning raw columns into feature vectors, computing statistics over billions of rows, finding similar items, mining frequent patterns and evaluating results honestly. Those tools are what this article covers.

How distributed training executes, including treeAggregate and tree histograms, is explained in Spark ML overview, and the Pipeline fit and transform contract and leak-free tuning in Spark ML pipelines. Here the focus is the toolkit: what each piece computes, what it costs on a cluster, a worked near-duplicate detection example, and how to migrate code still written against the old RDD-based API.

Advertisement

What MLlib means today

Spark ships two machine-learning packages under the MLlib name. org.apache.spark.ml (pyspark.ml in Python) is the DataFrame-based API and the one that receives new features. org.apache.spark.mllib (pyspark.mllib) is the original RDD-based API, in maintenance mode since Spark 2.0: it receives bug fixes but not new algorithms. The informal name spark.ml refers to the first package; MLlib is the umbrella for both.

In the DataFrame API, every algorithm reads its input from named columns and appends output columns. Features for estimators are one column of type Vector, dense or sparse. Most of the toolkit below is about producing that column well and about analysing data and predictions in it.

The MLlib toolkit around the estimators: features in, statistics, similarity, patterns, evaluation outRaw DataFramestrings, numbers, text, basketsFeature stagesindex, encode, impute, scale, hashfeatures: Vectordense or sparse, one columnStatisticsSummarizer, Correlation, ChiSquareSimilarity (LSH)MinHash, random projectionEstimatorsclassification, regression, ...Pattern miningFPGrowth, PrefixSpan on item arraysEvaluatorsAUC, F1, RMSE, silhouetteLegacy spark.mllibRDD API, maintenance onlyfit / transformpredictionsEverything on the left and centre is a DataFrame-based spark.ml API; the RDD API only receives fixes.
Feature stages turn raw columns into a single vector column. Statistics, LSH and estimators consume that column; pattern mining works directly on arrays of items; evaluators score prediction columns. The RDD-based API sits beside all of this and only receives fixes.

Feature transformers: which ones learn, and what they learn

Feature stages are either Transformers, which are stateless, or Estimators, which run a Spark job during fit to learn state and return a Model. Knowing which is which tells you what must be fitted on training data only and what costs a pass over the data.

StageLearns during fitNotes
StringIndexerlabel-to-index map, by frequency by defaulthandleInvalid: error, skip or keep (unseen go to an extra index)
OneHotEncodercategory count per columndropLast defaults to true; output is sparse
Imputermean or median (or mode) per columnmedian uses approximate quantiles
StandardScalerper-feature mean and standard deviationwithMean defaults to false, because centring densifies sparse vectors
QuantileDiscretizerbucket boundaries via approxQuantileproduces a Bucketizer; relativeError trades accuracy for speed
CountVectorizervocabulary of top termsvocabSize and minDF control size; needs a full pass
IDFdocument frequency per termpair with HashingTF or CountVectorizer
HashingTF, FeatureHashernothingstateless; columns come from a hash, so collisions are possible
Bucketizer, VectorAssembler, SQLTransformernothingpure functions of each row

Two rules follow. First, every learned stage must be fitted inside the Pipeline on training data only, or its statistics leak test information into the model. Second, every learned stage costs at least one job at fit time; a pipeline with ten StringIndexers used to mean ten passes, which is why Spark 3.0 added multi-column inputCols to StringIndexer and several other stages. Prefer the multi-column form.

Advertisement

Hashing versus vocabularies: collision arithmetic

Text and high-cardinality categorical features can be vectorised with a learned vocabulary (CountVectorizer, StringIndexer plus OneHotEncoder) or with the hashing trick (HashingTF, FeatureHasher), which maps each term to hash(term) mod numFeatures using MurmurHash3. Hashing needs no fit, no vocabulary to store and handles unseen terms gracefully, but distinct terms can share a column.

The arithmetic is worth doing. With n distinct terms hashed into m buckets, the expected number of occupied buckets is m(1 - e^(-n/m)). The default numFeatures is 2^18 = 262,144. For 100,000 distinct terms, n/m is about 0.38, so about 83,000 buckets are occupied and roughly 17,000 terms share a column with another term. Raising numFeatures to 2^20 cuts that to about 4,600, around 4.6% of terms, at the cost of wider sparse vectors and, for linear models, a larger coefficient vector to aggregate each iteration. Use a power of two, since the modulo otherwise distributes terms unevenly.

Choose vocabularies when you need interpretability (mapping a coefficient back to a word) or when the vocabulary is stable and small; choose hashing for streaming or for very large, evolving vocabularies, sized from the calculation above.

Distributed statistics

The statistics package computes summaries in one pass using mergeable aggregators, so the cost is a scan plus a small exchange. Summarizer computes per-feature mean, variance, min, max, count, non-zero count and norms on a vector column. Correlation.corr returns a Pearson or Spearman correlation matrix; Spearman must rank every column, which involves sorting and is far more expensive than Pearson. ChiSquareTest tests independence between each categorical feature and a categorical label. For scalar columns, DataFrame.approxQuantile uses a Greenwald-Khanna style sketch with a configurable relative error; a relative error of 0 computes exact quantiles and is expensive.

from pyspark.ml.stat import Summarizer, Correlation, ChiSquareTest
from pyspark.sql import functions as F

summary = (df.select(Summarizer.metrics("mean", "variance", "numNonZeros", "max")
                     .summary(F.col("features")).alias("s"))
             .select("s.*").first())

pearson = Correlation.corr(df, "features", "pearson").head()[0]   # a DenseMatrix on the driver
chi = ChiSquareTest.test(df, "cat_features", "label").head()       # pValues, degreesOfFreedom, statistics

p50, p99 = df.approxQuantile("latency_ms", [0.5, 0.99], 0.001)

Watch the driver. A correlation matrix over 20,000 features has 400 million entries and is collected to the driver as a dense matrix, which will not fit. Select features first, or compute correlations on a sample and on the columns you care about.

Similarity search with locality-sensitive hashing

Finding pairs of similar items among millions by comparing every pair is quadratic. Locality-sensitive hashing (LSH) hashes items so that similar ones land in the same bucket with high probability, then only compares items that share a bucket. MLlib provides MinHashLSH for Jaccard distance on sets represented as sparse binary vectors, and BucketedRandomProjectionLSH for Euclidean distance. Both offer approxSimilarityJoin, which returns pairs within a distance threshold, and approxNearestNeighbors for the k nearest items to one key.

For MinHash, the probability that two sets collide in one hash table equals their Jaccard similarity s. Spark combines numHashTables tables with OR-amplification: a pair becomes a candidate if it collides in any table, with probability 1 - (1 - s)^L. With L = 5, a pair with similarity 0.8 is found with probability 0.9997, while a pair at 0.2 still becomes a candidate with probability 0.67. Those false candidates are removed by the exact distance check, but they cost shuffle and compute. More tables improve recall and increase cost; the threshold decides precision.

Worked example: near-duplicate product listings. A marketplace has 30 million listings and wants titles that are near-copies. Tokenise titles into character 3-shingles, hash them into a large binary vector, and self-join with MinHash:

from pyspark.ml import Pipeline
from pyspark.ml.feature import RegexTokenizer, NGram, HashingTF, MinHashLSH
from pyspark.sql import functions as F

prep = Pipeline(stages=[
    RegexTokenizer(inputCol="title", outputCol="chars", pattern="", minTokenLength=1),
    NGram(n=3, inputCol="chars", outputCol="shingles"),
    HashingTF(inputCol="shingles", outputCol="vec", numFeatures=1 << 20, binary=True),
]).fit(listings)
vec = (prep.transform(listings)
           .filter(F.size("shingles") > 0)          # MinHash rejects all-zero vectors
           .select("listing_id", "vec").cache())

lsh = MinHashLSH(inputCol="vec", outputCol="hashes", numHashTables=5).fit(vec)
pairs = (lsh.approxSimilarityJoin(vec, vec, threshold=0.3, distCol="jaccard_dist")
            .filter("datasetA.listing_id < datasetB.listing_id")
            .select(F.col("datasetA.listing_id").alias("a"),
                    F.col("datasetB.listing_id").alias("b"), "jaccard_dist"))

A threshold of 0.3 Jaccard distance means at least 70% shingle overlap. Two practical points: popular shingles create huge buckets that dominate the join, so remove boilerplate words or very common shingles first; and run the self-join on a sample to measure candidate volume before launching it on the full table.

Frequent patterns: FPGrowth and PrefixSpan

FPGrowth mines frequent itemsets from an array column, such as the products in each basket, without generating candidate sets explicitly. Fitting returns freqItemsets and associationRules with confidence, lift and support, and transform predicts items to add for each basket from the rules. PrefixSpan mines frequent sequential patterns, such as page paths, and returns them from findFrequentSequentialPatterns.

from pyspark.ml.fpm import FPGrowth

baskets = orders.groupBy("order_id").agg(F.collect_set("sku").alias("items"))
fp = FPGrowth(itemsCol="items", minSupport=0.001, minConfidence=0.2, numPartitions=400)
model = fp.fit(baskets)
rules = model.associationRules.filter("lift > 3").orderBy(F.desc("lift"))

Support is the dangerous knob. The number of frequent itemsets grows explosively as minSupport falls, and a threshold that is ten times too low can exhaust executor memory. Start high, count results, and lower it step by step. Items must be unique within a row, which is why the example uses collect_set.

Evaluation without fooling yourself

MLlib's evaluators compute metrics as distributed aggregations: BinaryClassificationEvaluator (areaUnderROC, areaUnderPR) from the rawPrediction column, MulticlassClassificationEvaluator (f1 by default, accuracy, weighted precision and recall), RegressionEvaluator (rmse by default, mae, r2), ClusteringEvaluator (silhouette) and RankingEvaluator for recommendations. They plug into CrossValidator and TrainValidationSplit.

Three traps are common. Area under ROC looks excellent on heavily imbalanced data even when the model is useless at the operating point you need; use areaUnderPR and precision at a chosen threshold. Evaluating a random split of time-ordered or user-grouped data leaks the future or the same user into the test set; split by time or by entity. And a single metric hides slices; group predictions by segment and compute the metric per segment with ordinary DataFrame aggregations, as below.

from pyspark.ml.evaluation import BinaryClassificationEvaluator

ev = BinaryClassificationEvaluator(labelCol="label", rawPredictionCol="rawPrediction",
                                   metricName="areaUnderPR")
print("overall PR AUC", ev.evaluate(preds))

# Per-segment metrics: evaluate each slice, and look at precision at the threshold you will ship
for seg in [r["segment"] for r in preds.select("segment").distinct().collect()]:
    part = preds.filter(F.col("segment") == seg)
    tp = part.filter("prediction = 1 AND label = 1").count()
    fp = part.filter("prediction = 1 AND label = 0").count()
    print(seg, round(ev.evaluate(part), 3), "precision", tp / max(tp + fp, 1))

The loop runs one job per segment, which is fine for a handful of segments; for hundreds, compute the confusion counts in a single groupBy on segment, label and prediction instead. Record the evaluation split, the metric and the threshold alongside the persisted model, so a later retrain is compared against the same yardstick rather than against a number nobody can reproduce.

Leaving the RDD-based spark.mllib API

Older code uses LabeledPoint, pyspark.mllib.linalg.Vectors and algorithms such as LogisticRegressionWithLBFGS on RDDs. Migrating brings Catalyst-optimised input handling, pipelines and persistence, and it is a precondition for Spark Connect: Connect does not support the RDD API at all, while Spark 4.0 can run the regular pyspark.ml estimators over a Connect session (earlier versions support far less, so check your version). The vector types of the two packages are different classes, so a DataFrame built from old vectors must be converted:

from pyspark.mllib.util import MLUtils
from pyspark.ml.classification import LogisticRegression

# Before: RDD[LabeledPoint] -> DataFrame with old-style vectors
df_old = labeled_points_rdd.map(lambda lp: (float(lp.label), lp.features)).toDF(["label", "features"])
df = MLUtils.convertVectorColumnsToML(df_old, "features")     # now pyspark.ml.linalg vectors

model = LogisticRegression(maxIter=100, regParam=0.01).fit(df)

Migrate one model at a time and compare coefficients and metrics against the old implementation on the same data; defaults for regularisation, standardisation and intercepts differ between some old and new algorithms. Persist the new model with model.write().overwrite().save(path) and replace any custom serialisation. Caching guidance from Spark persistence still applies to the iterative estimators.

Failure modes

  • Unseen categories at scoring time. StringIndexer with the default handleInvalid raises an error on a new value in production. Use keep, and monitor how often the extra index appears.
  • Densified vectors. Setting withMean=True on StandardScaler, or a stage that converts to dense, turns a 1-million-wide sparse vector into a dense one and runs executors out of memory.
  • Driver collection. Correlation matrices, large vocabularies and huge association-rule sets end up on the driver. Size them before collecting.
  • Skewed LSH buckets. A handful of common shingles or zero-heavy vectors create giant buckets and straggler tasks.
  • Silent hash collisions. numFeatures chosen once and never revisited as the vocabulary grows. Recompute the collision estimate when the data changes.
  • Mixed vector classes. Passing mllib vectors to ml estimators fails with a type error far from the cause; convert at the boundary.

What to do next

  1. Inventory the feature stages in one production pipeline, mark which ones learn state, and confirm each is fitted only on training data.
  2. Switch repeated single-column StringIndexer or OneHotEncoder stages to the multi-column form and measure fit time before and after.
  3. Compute the hashing collision estimate for every HashingTF or FeatureHasher you use, and resize numFeatures where more than a few percent of terms collide.
  4. Run Summarizer and approxQuantile over your feature vectors and label to catch empty, constant or extreme features before training.
  5. Try MinHashLSH on a 1% sample of a deduplication or similar-items problem, measuring candidate counts and precision before scaling up.
  6. List remaining pyspark.mllib usages and migrate the highest-traffic model first, comparing metrics with the old version.
Key takeaway: MLlib is more than its estimators. Feature stages either learn state, and must be fitted on training data inside a pipeline, or are pure functions; hashing trades collisions for statelessness, and the collision rate can be calculated. The statistics package summarises vectors in one pass, LSH makes similarity joins practical at scale, FPGrowth and PrefixSpan mine patterns if support is chosen carefully, and evaluators need imbalance-aware metrics and honest splits. The RDD-based API is in maintenance mode, so migrate to the DataFrame API one model at a time.