A model that scores well on the data it was tuned on tells you very little about how it will score next month. Cross-validation is the standard way to get a better estimate without spending a separate validation set on every decision: split the data into k folds, train on k minus 1 of them, score on the one left out, rotate, and average. Spark ML packages this as CrossValidator in pyspark.ml.tuning.

The everyday usage, wrapping the whole pipeline so preprocessing is refitted per fold and budgeting the number of fits, is covered in Spark ML Pipelines, in depth. This article goes underneath. It explains exactly what CrossValidator does with your DataFrame, how to design folds when rows are grouped or ordered in time, a nondeterminism trap that silently corrupts folds, how to read avgMetrics and stdMetrics instead of trusting the argmax, what the work costs on a cluster, and when to write your own loop. Examples use PySpark; version notes are given where a feature is recent.

Advertisement

What cross-validation estimates, and what it does not

You have a procedure, for example: index these columns, assemble features, fit logistic regression with regParam 0.01, and you want to know how well its models will score on new data. One hold-out split gives one noisy answer. k-fold cross-validation gives k answers, each from a model trained on a fraction (k-1)/k of the data and scored on rows it never saw, and averages them. The average is less noisy than a single split, and the spread across folds tells you how noisy it is.

Two uses are often confused. Model selection compares procedures, such as nine settings of regularisation, and picks one. Model assessment reports how good the chosen one is. If you use the same cross-validation score for both, the reported number is optimistic: you picked the setting that got lucky on these folds, the winner's curse. The fix is cheap on big data: keep an untouched test set, run cross-validation on the rest, and score the final model on the test set once.

Everything rests on one assumption: each validation fold looks like the future data relative to its training folds. Random splitting makes that true when rows are independent. It is false when the same customer, device or patient appears in many rows, and false when the target depends on time.

What CrossValidator does, step by step

The PySpark implementation is short and worth knowing, because every failure mode below follows from it. Given an estimator, a list of m param maps, an evaluator, numFolds (default 3), a seed, parallelism (default 1), collectSubModels (default false) and optionally foldCol (Spark 3.1 and later), fit does the following.

  1. Assign folds. Without foldCol, it adds a column rand(seed) and, for fold i, takes rows whose value falls in [i/k, (i+1)/k) as validation and the rest as training. With foldCol, it uses your integer column, which must lie in [0, k), and no random split happens. The Scala implementation uses a seeded per-fold sampler over the underlying RDD; the idea is the same.
  2. Cache the fold. For each fold in turn, the training and validation DataFrames are persisted.
  3. Fit and score in a thread pool. One task per param map is submitted to a driver-side pool of size min(parallelism, m). Each task fits the estimator on the training DataFrame through fitMultiple, transforms the validation DataFrame and calls the evaluator. Every fit is a full set of Spark jobs.
  4. Release the fold. Both DataFrames are unpersisted before the next fold starts, so folds run in sequence and parallelism exists only across param maps.
  5. Aggregate. The k by m metric matrix is averaged per param map into avgMetrics; since Spark 3.3 the per-map standard deviation is exposed as stdMetrics.
  6. Select and refit. The best index is the argmax if evaluator.isLargerBetter(), otherwise the argmin. The estimator is fitted once more with that param map on the entire input, and the result becomes bestModel.
What CrossValidator.fit does with one DataFrame, k folds and m param mapsInput DataFramedev set onlyFold assignmentrand(seed) or foldColfold 1: train + validfold 2: train + valid... fold k (cached, one at a time)Driver thread poolmin(parallelism, m) fitsper foldEvaluatormetric per fitk x m metric matrixmean and std per mapSelect best mapargmax or argminRefit on all inputbestModelUntouched test setevaluated once, by you, outside CrossValidatorfinal checkFolds run in sequence.Param maps within a foldrun concurrently.k*m + 1 full fits
CrossValidator in PySpark. The only randomness is the fold assignment; everything else is deterministic bookkeeping around k*m + 1 ordinary fits.

Two consequences follow. The refit uses all of the data you passed in, so a test set must be removed before fit. And the selection rule is a bare argmax, however small the margin.

Advertisement

A worked example: churn with group-aware folds

Suppose a table of customer-month rows with features and a churned label, about 0.8% positive. Customers appear in many months, so random row splits would put the same customer in training and validation. The example removes a test set by customer, builds deterministic group-aware folds, tunes nine settings and reads the results.

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

spark = SparkSession.builder.appName("churn-cv").getOrCreate()
df = spark.read.parquet("s3://example-bucket/churn/features/")

# 1. Test set by customer, removed before any tuning.
df = df.withColumn("bucket", F.pmod(F.xxhash64("customer_id"), F.lit(10)))
test = df.filter("bucket = 0")
dev = df.filter("bucket != 0")

# 2. Deterministic folds from a SALTED hash of the group key.
K = 5
dev = dev.withColumn(
    "fold", F.pmod(F.xxhash64(F.lit("cv-salt"), "customer_id"), F.lit(K)).cast("int"))
dev = dev.persist()
dev.groupBy("fold").agg(F.count("*").alias("rows"), F.avg("churned").alias("pos_rate")).show()

# 3. Pipeline and grid: 3 x 3 = 9 param maps.
plan_idx = StringIndexer(inputCol="plan", outputCol="plan_idx", handleInvalid="keep")
plan_ohe = OneHotEncoder(inputCols=["plan_idx"], outputCols=["plan_vec"])
assemble = VectorAssembler(inputCols=["plan_vec", "tenure_months", "tickets_90d", "spend_90d"],
                           outputCol="features")
lr = LogisticRegression(labelCol="churned", featuresCol="features", maxIter=50)
pipe = Pipeline(stages=[plan_idx, plan_ohe, assemble, lr])

grid = (ParamGridBuilder()
        .addGrid(lr.regParam, [0.001, 0.01, 0.1])
        .addGrid(lr.elasticNetParam, [0.0, 0.5, 1.0])
        .build())
pr_auc = BinaryClassificationEvaluator(labelCol="churned", metricName="areaUnderPR")

cv = CrossValidator(estimator=pipe, estimatorParamMaps=grid, evaluator=pr_auc,
                    numFolds=K, foldCol="fold", parallelism=3)
cv_model = cv.fit(dev)                      # 9 x 5 = 45 pipeline fits + 1 refit

for pm, mean, std in zip(grid, cv_model.avgMetrics, cv_model.stdMetrics):
    print(pm[lr.regParam], pm[lr.elasticNetParam], round(mean, 4), round(std, 4))

print("test PR AUC:", pr_auc.evaluate(cv_model.bestModel.transform(test)))

Why the salt? The test split used xxhash64(customer_id) mod 10. If the folds used the same hash mod 5, then among the remaining customers (hash mod 10 not equal to 0), fold 0 would receive only those with hash mod 10 equal to 5, half the share of every other fold. Two splits derived from one hash are correlated; salting the second makes them independent. The groupBy("fold") check is how you would notice: rows and positive rate per fold should be roughly equal.

The fold column travels through the pipeline untouched because VectorAssembler reads only the columns it is named; never let a stage select all numeric columns, or the fold number becomes a leaking feature.

Fold design: random, grouped, stratified and time-ordered

Random rows are fine when rows are independent draws: one row per customer, per image, per document. That is the default rand(seed) path.

Grouped data needs every row of a group in the same fold, otherwise the model memorises the group in training and recognises it in validation, and scores look better than they will be on new customers. Hash the group key into foldCol as above. Watch for heavy groups that make one fold far larger than the rest.

Rare positives make fold metrics noisy, and Spark does not stratify. You can build a stratified foldCol by numbering rows within each label in hashed order, but when data is also grouped, grouping matters more; check per-fold positive rates instead.

Time-ordered targets cannot use CrossValidator at all, even with foldCol set to the month, because it trains on every other fold, including later ones, and so learns from the future to predict the past. Use forward chaining: train on everything before a cutoff, validate on the window after it, move the cutoff. fitMultiple gives you the same per-param-map fitting the built-in class uses:

import numpy as np

def forward_chain_cv(est, grid, evaluator, df, time_col, cutoffs, gap_days=0):
    """Fold i trains on [start, cutoffs[i]) and validates on [cutoffs[i] + gap, cutoffs[i+1])."""
    scores = np.zeros((len(cutoffs) - 1, len(grid)))
    for i in range(len(cutoffs) - 1):
        train = df.filter(F.col(time_col) < F.lit(cutoffs[i])).persist()
        valid = df.filter((F.col(time_col) >= F.date_add(F.lit(cutoffs[i]), gap_days)) &
                          (F.col(time_col) < F.lit(cutoffs[i + 1]))).persist()
        for j, model in est.fitMultiple(train, grid):
            scores[i, j] = evaluator.evaluate(model.transform(valid))
        train.unpersist()
        valid.unpersist()
    return scores.mean(axis=0), scores.std(axis=0), scores

The gap matters when the label looks forward. A churn label meaning churned within 30 days makes rows in the last 30 days before a cutoff carry information about the validation window; dropping that band from validation, or from training, removes the overlap. Keep the full matrix: a setting that wins early and loses recently signals drift.

The nondeterministic-input trap

The default split looks safe because it is seeded. It is not, if the input is not stable. rand(seed) gives each row a value that depends on the partition it is in and its position within that partition. CrossValidator builds the training set and the validation set as two separate filters over the same lineage, and each is computed when it is cached. If computing the input twice can produce rows in a different order or a different partitioning, a row can draw a different random value each time and end up in both training and validation, or in neither.

Typical causes are a repartition(n) without columns, which distributes rows round-robin, a union of sources read in parallel, or a partition recomputed after an executor is lost. Nothing fails. Scores just come out a little too good, and differ between reruns with the same seed.

Two fixes. Persist and materialise the input before calling fit, preferably with a storage level that spills to disk so eviction does not trigger recomputation, or checkpoint it. Better, assign folds from a hash of a stable key with foldCol, which is deterministic by construction and also gives you group awareness and reproducibility across jobs and Spark versions.

Reading avgMetrics and stdMetrics

avgMetrics is in the same order as getEstimatorParamMaps(). Print them side by side, sorted, every time. The table below shows the shape of a typical result for the example; the numbers are illustrative, not a benchmark.

regParamelasticNetmean PR AUCstd across folds
0.0010.00.3120.021
0.010.00.3180.019
0.010.50.3160.020
0.10.00.3090.018
0.11.00.2740.025

The argmax picks regParam 0.01 with no elastic net. But the top four settings differ by less than half a standard deviation, and the standard error of a k-fold mean is roughly the standard deviation divided by the square root of k. The honest reading is that they are tied and the heavily sparse model is worse. A common rule is the one-standard-error rule: pick the simplest, most regularised setting whose mean is within one standard error of the best. CrossValidator will not do that for you, but refitting is one line:

import math
best = int(np.argmax(cv_model.avgMetrics))          # use argmin for RMSE-style metrics
se = cv_model.stdMetrics[best] / math.sqrt(K)
ok = [i for i, m in enumerate(cv_model.avgMetrics) if m >= cv_model.avgMetrics[best] - se]
chosen = max(ok, key=lambda i: grid[i][lr.regParam])  # most regularised among the tied
final_model = pipe.fit(dev, grid[chosen])

Set collectSubModels=True only for diagnosis; it keeps k times m fitted models referenced from the driver.

Cost, parallelism and when to use TrainValidationSplit

Every one of the k*m + 1 fits runs the whole pipeline; logistic regression with maxIter=50 can make up to 50 distributed passes per fit. Three levers reduce the work.

  • Hoist stages that learn nothing. SQL feature engineering and VectorAssembler fit no statistics, so apply them once, cache the result and cross-validate only the stages that learn. Stages that learn from data, such as imputers, scalers and indexers, belong inside, or their statistics leak across folds.
  • Search less. Use log-spaced values, tune coarse then fine, or sample a random subset of a larger grid with random.sample(full_grid, 20).
  • Split once. TrainValidationSplit (default trainRatio 0.75) costs m + 1 fits and is often precise enough on large data.

parallelism runs several fits of the same fold at once from driver threads. It helps when one fit cannot keep the cluster busy and hurts when fits already saturate executors, since each concurrent fit holds its own caches. Start at 2 to 4, enable spark.scheduler.mode=FAIR so concurrent jobs share executors, and watch task timelines and storage in the Spark UI; Spark memory management explains what spills when you overdo it.

Failure modes

  • Leaky preprocessing. A scaler or imputer fitted on the full data before cross-validation; validation folds influenced their own features.
  • Group leakage. Random row folds over repeated customers or sessions; the scores collapse in production on new entities.
  • Future leakage. Built-in folds over time-dependent targets, or forward-chaining without a gap when labels look ahead.
  • Overlapping folds. The default random split over an unstable, recomputed input; inflated and non-reproducible scores.

Trade-offs

More folds buy a less biased estimate and a visible spread at linear cost in fits; a single split is cheaper and often enough at scale. Forward chaining respects time but gives early folds less training data. Selecting by argmax is simple and reproducible; the one-standard-error rule trades a tiny expected loss for simpler, more stable models. Hoisting stages saves a lot of compute but must be limited to stages that learn nothing. For choosing storage levels when you cache fold inputs, see Spark persistence.

What to do next

  1. Remove a test set before tuning, split by the same unit you will predict for, and never pass it to fit.
  2. Decide whether rows are independent, grouped or time-ordered, and build a deterministic foldCol from a salted hash, or a forward-chaining loop for time.
  3. Choose the evaluator metric for the business question, and confirm isLargerBetter for custom ones.
  4. Hoist non-learning stages, cache their output, and size the grid so k*m + 1 fits fit your budget.
  5. Print avgMetrics with stdMetrics for every run, apply a one-standard-error rule, and refit the chosen setting.
  6. Score the final model once on the test set and record that number, the seed, the fold definition and the grid with the model.
Key takeaway: CrossValidator is a small loop: assign folds, cache one fold at a time, fit every param map in a driver thread pool, average, take the argmax and refit on everything. Its estimate is only as good as its folds, so build them deterministically from a salted hash of the unit you predict for, switch to forward chaining when time matters, read the spread before trusting the winner, and keep a test set that tuning never touches.