Spark ML ships nine classifiers behind one interface. You choose between them, and tune them, by understanding three things: how each one trains on a cluster, what the three output columns it writes actually mean, and how the decision threshold turns a score into a label. Most classification bugs in Spark are not algorithm bugs. They come from treating a raw score as a probability, a default threshold of 0.5 that was never right for the business, or class weights that quietly shift every predicted probability.

This article covers the classifier catalog as of the Spark 4.2 documentation, the training mechanics behind it, the output columns, thresholds, class imbalance, evaluation and calibration, and works through an imbalanced fraud model. Pipelines and hyperparameter search have their own articles: Spark ML Pipelines and Spark ML cross-validation. For a survey of the whole library, start with Spark ML Overview.

The classifier catalog

Every classifier is an Estimator. It reads a vector featuresCol and a numeric labelCol (labels 0, 1, ..., K−1) and returns a Model that is a Transformer. The table summarises how they differ.

ClassifierClassesHow it trainsUse when
LogisticRegressionBinary, or multiclass with family="multinomial"Iterative quasi-Newton over the full dataStrong, interpretable baseline; calibrated-ish probabilities
LinearSVCBinary onlyIterative optimisation of hinge lossLinear margin classifier; no probability output
DecisionTreeClassifierBinary and multiclassBinned split search, level by levelExplainability; base learner for ensembles
RandomForestClassifierBinary and multiclassMany trees trained in parallel on bootstrap samplesRobust default for tabular data
GBTClassifierBinary onlyTrees trained sequentially on gradientsHigher accuracy on tabular data, at the cost of training time
NaiveBayesBinary and multiclassOne pass of countingText and count features; very fast
MultilayerPerceptronClassifierBinary and multiclassFeed-forward net trained with L-BFGSSmall dense networks; not a deep learning framework
FMClassifierBinaryFactorisation machine with pairwise interactionsSparse, high-cardinality features
OneVsRestMulticlass via a binary baseK binary models, one per classUsing a binary-only classifier on K classes

The Spark guide documents LinearSVC as binary only. GBTClassifier also rejects labels other than 0 and 1, so for multiclass gradient boosting you either wrap it in OneVsRest or use an external library. Note the documented detail on multinomial logistic regression: with family="multinomial" the model exposes coefficientMatrix and interceptVector rather than coefficients and intercept, even when it is used for binary classification.

How training runs on a cluster

Driverweights w, L-BFGS stateExecutor 1cached partitionsExecutor 2cached partitionsExecutor ncached partitionsbroadcast wtreeAggregatesum of loss and gradientback todriverOne iteration = one pass over the cached data. maxIter iterations = maxIter passes.
Figure 1. One iteration of distributed logistic regression. The driver broadcasts the weights, each executor computes loss and gradient over its cached partitions, and a tree aggregation sums them back to the driver for the optimiser step.

Linear models. LogisticRegression and LinearSVC are iterative. Each iteration broadcasts the current weights, every executor computes the loss and gradient over its partitions, and treeAggregate combines the partial sums before they reach the driver. For logistic regression the optimiser is L-BFGS for L2 or no regularisation, and OWL-QN when an L1 component is present. The cost model is simple: one pass over the data per iteration, up to maxIter passes. Cache the training DataFrame, or every iteration re-reads and re-parses the source. The aggregationDepth parameter adds levels to the aggregation tree when there are many partitions or very wide feature vectors.

Trees. Spark does not sort continuous features. It discretises each one into at most maxBins candidate thresholds, default 32, from a sample of quantiles, then grows trees level by level. Executors compute split statistics for every active node, and the driver chooses the best splits. maxMemoryInMB bounds how many nodes are processed per pass. Random forests grow all their trees in these same passes. GBT cannot, because each tree fits the gradients of the ensemble so far, so maxIter trees mean that many sequential rounds. Set a checkpoint directory and checkpointInterval for long GBT runs, or the lineage grows with every round.

NaiveBayes is a single aggregation pass. MultilayerPerceptronClassifier uses the same iterative pattern as the linear models, with a larger gradient.

rawPrediction, probability and prediction

A fitted classification model typically writes three columns, and they mean different things:

  • rawPrediction: a vector of per-class confidence scores on the model's own scale. For logistic regression these are margins, for tree ensembles vote totals or summed tree predictions, and for LinearSVC the signed distance from the hyperplane. Raw scores rank well, but they are not probabilities.
  • probability: a vector of class probabilities, written by models that can produce one. LinearSVC does not produce one at all. Some models produce probabilities that rank well but are poorly calibrated.
  • prediction: the chosen label after the threshold is applied.

For binary logistic regression the relationship is exact. With margin m = w·x + b, rawPrediction is [−m, m] and probability is [1 − σ(m), σ(m)], where σ is the logistic function. A margin of 1.2 gives a probability of about 0.77 for class 1. Thresholding at 0.5 on probability is the same as thresholding at 0 on the margin, which is why thresholding the raw column at 0.5 is a silent bug.

For binary models, threshold (default 0.5) is compared with the probability of class 1. For multiclass, thresholds is an array with one value per class, and the predicted class is the one with the largest p/t. Setting a class's threshold lower makes it easier to predict. On binary LogisticRegression the two parameters must agree if both are set, so set only one. Changing the threshold never requires retraining. It is a parameter of the fitted model, so you can pick it after training, from a validation set.

Class imbalance and weightCol

With 0.2 percent positives, a model can be 99.8 percent accurate by predicting no fraud for everything. There are three standard responses. They can be combined, but each changes the probabilities differently.

  • Class weights through weightCol. LogisticRegression, LinearSVC, NaiveBayes and the tree classifiers accept a weight column. Weighting positives by w scales their contribution to the loss.
  • Downsampling the majority class. This is cheaper to train, and the parallelism is no longer spent mostly on negatives. Use sampleBy for a stratified sample.
  • Doing nothing to the data and moving the threshold. This is often enough for models that rank well. Evaluate it first, because it is the only option that leaves probabilities untouched.
from pyspark.sql import functions as F

pos_rate = train.agg(F.avg("label")).first()[0]
w_pos = (1 - pos_rate) / pos_rate             # balance total weight of the two classes
train_w = train.withColumn(
    "w", F.when(F.col("label") == 1, F.lit(w_pos)).otherwise(F.lit(1.0)))

from pyspark.ml.classification import LogisticRegression
lr = LogisticRegression(featuresCol="features", labelCol="label",
                        weightCol="w", regParam=0.01, elasticNetParam=0.0)
model = lr.fit(train_w.cache())

Both weighting and downsampling inflate predicted probabilities. If positives were weighted by w, or negatives downsampled by a factor w, the model's odds are about w times too high. Correct them with oddstrue = oddsmodel / w before using probabilities as probabilities, for example in an expected-loss calculation. The ranking, and therefore ROC and PR curves, is unaffected.

Evaluation and choosing a threshold

BinaryClassificationEvaluator reads rawPredictionCol and offers areaUnderROC (the default) and areaUnderPR. Under heavy imbalance, prefer area under PR: ROC is dominated by the huge negative class and looks excellent even for models that are useless at the operating point. MulticlassClassificationEvaluator reads prediction and offers metrics including f1 (the default), accuracy, weightedPrecision, weightedRecall, per-label metrics such as precisionByLabel together with metricLabel, and logLoss, which reads the probability column. Both evaluators accept a weight column.

Neither area metric tells you which threshold to ship. The operating point comes from costs. If a missed fraud costs 200 and a false alarm costs 5 in review time, sweep thresholds on a validation set and minimise the total cost:

from pyspark.ml.functions import vector_to_array

scored = (model.transform(valid)
          .withColumn("p1", vector_to_array("probability")[1])
          .withColumn("bucket", F.floor(F.col("p1") * 1000) / 1000)
          .groupBy("bucket")
          .agg(F.sum("label").alias("pos"), F.count("*").alias("n"))
          .toPandas().sort_values("bucket", ascending=False))

scored["tp"] = scored["pos"].cumsum()                 # predicted positive at threshold = bucket
scored["fp"] = (scored["n"] - scored["pos"]).cumsum()
total_pos = scored["pos"].sum()
scored["cost"] = 200 * (total_pos - scored["tp"]) + 5 * scored["fp"]
best = scored.loc[scored["cost"].idxmin(), "bucket"]
model.setThreshold(float(best))

Bucketing at 0.001 keeps the collected result to at most a thousand rows, however large the validation set is. The threshold is chosen on probabilities from the weighted model, so recompute it if you change the weights.

Calibration

Logistic regression without weighting tends to be reasonably calibrated. Random forest probabilities are averaged votes, and they are usually pulled away from 0 and 1. GBT probabilities come from a logistic transform of the boosted margin, and they can be overconfident. Check before trusting any of them. Bucket validation predictions into ten bins by probability and compare the mean prediction with the observed positive rate in each bin. If they disagree and you need real probabilities, fit a one-feature logistic regression on the raw score using held-out data, a form of Platt scaling, and apply it as a final pipeline stage. The binning check is a few lines:

calib = (model.transform(valid)
         .withColumn("p1", vector_to_array("probability")[1])
         .withColumn("bin", F.least(F.floor(F.col("p1") * 10), F.lit(9)))
         .groupBy("bin")
         .agg(F.avg("p1").alias("predicted"), F.avg("label").alias("observed"),
              F.count("*").alias("n"))
         .orderBy("bin"))
calib.show()

Worked example: an imbalanced fraud model

Take 50 million card transactions with 0.2 percent fraud and 40 features, some of them high-cardinality merchant categories. The goal is a model that flags transactions for review, under a review budget.

  1. Split by time, training on earlier months and validating on the latest. A random split leaks fraud patterns across the boundary and overstates PR.
  2. Baseline: LogisticRegression with weights as above, one-hot merchant category and a standardised amount. Cache the assembled features. At 100 iterations that is 100 passes, so check in the Spark UI that each pass reads from memory.
  3. Challenger: GBTClassifier with maxDepth 5, maxIter 100 and checkpointing, trained on all positives plus a 10 percent sample of negatives for speed. Remember that this multiplies the odds by 10.
  4. Compare areaUnderPR on the time-ordered validation month, then sweep each model's threshold against the cost function and the review budget, for example at most 2,000 alerts a day.
  5. Calibrate the winner if downstream systems consume probabilities, then fix the threshold in the saved PipelineModel and record it with the model version.

In practice, the gain from GBT over a well-featured logistic regression on data like this is often smaller than the gain from better features and a correct threshold. Measure it rather than assuming it.

Failure modes

  • Treating rawPrediction as a probability, for example thresholding margins at 0.5.
  • Shipping the default threshold of 0.5 on imbalanced data, so almost nothing is flagged.
  • Using weighted probabilities as real risks without the odds correction.
  • Uncached training data for iterative models, so every iteration re-reads the source.
  • Labels that are not 0 to K−1, such as strings or 1 and 2. Index them with StringIndexer first.
  • Too few bins for a feature that matters at fine resolution, and categorical features with more categories than maxBins, which fail at fit time.
  • Long GBT lineage without checkpointing, which slows each round and can overflow the stack.
  • ROC-only evaluation under imbalance, which hides a poor operating point.

Trade-offs

Linear versus trees. Linear models are fast, stable and interpretable, but they need engineered interactions. Trees find interactions on their own, need less preprocessing and cost more passes. Random forest versus GBT: forests parallelise across trees and are hard to overfit, while GBT is usually more accurate and is strictly sequential. Weights versus sampling: weights keep all the data, while sampling cuts cost and also loses information about negatives. Spark ML versus single-node libraries: if the training set fits on one large machine, a single-node gradient boosting library is often faster and richer. Spark's advantage is that training and scoring stay next to data that does not fit there. OneVsRest versus native multiclass: OneVsRest trains K models, so K times the passes, and its per-class scores are not jointly normalised. Prefer multinomial logistic regression or a native multiclass tree model when one fits.

What to do next

  1. Pick a baseline: weighted LogisticRegression on cached, assembled features.
  2. Check which classifiers are binary only before designing a multiclass pipeline.
  3. Evaluate with areaUnderPR under imbalance and a time-ordered split where time matters.
  4. Choose the threshold from a cost sweep on validation data, then setThreshold on the fitted model.
  5. Correct odds for any weighting or downsampling before using probabilities as risks.
  6. Run a ten-bin calibration check and add Platt scaling if probabilities are consumed downstream.
  7. Checkpoint long GBT runs and size maxBins to your categorical features.
  8. Tune hyperparameters with cross-validation only after the threshold and weighting are settled.
Key takeaway: Spark ML classifiers share one contract: a features vector and 0-based labels in, rawPrediction, probability and prediction out. Know which models are binary only, cache data for iterative learners, evaluate with area under PR when classes are imbalanced, pick the threshold from costs on validation data, correct probabilities for any weighting or sampling, and check calibration before treating scores as risks.