
Cross-Fitting and Out-of-Fold Prediction
Source:vignettes/articles/Crossfitting-and-OOF-Prediction.Rmd
Crossfitting-and-OOF-Prediction.RmdWhy cross-fit?
super_learner() already uses cross-validation
internally: each candidate learner is trained on each fold’s training
split and scored on its held-out split, and the ensemble weights are
chosen to minimize held-out loss. So the learners never score
their own training data. But the ensemble weights are estimated
from the pooled held-out predictions of every observation — including
observation
’s.
For most predictive uses this is fine. For semiparametric causal
inference (AIPW, TMLE, double machine learning), the theory asks for a
little more: each observation’s nuisance prediction
should come from a model fit entirely without observation
— weights included.
crossfit_super_learner() provides exactly that. It
splits the data into n_folds outer folds and fits a
complete super_learner() — candidate learners and
ensemble weights, via its own inner_n_folds-fold
cross-validation — on each outer training split, then predicts on the
corresponding held-out fold. Every out-of-fold prediction therefore
comes from an ensemble whose base learners and meta-learned
weights never saw that row.
The three fitting functions answer three different questions:
-
super_learner(): “give me the best predictor of ” — use$predict()on new data; -
crossfit_super_learner(): “give me honest per-observation predictions for my training rows” — there deliberately is no single$predict()(calling it errors with directions), because a cross-fitted model is its collection of fold-specific fits; -
cv_super_learner(): “how good is the super learning procedure itself?” — an honest estimate of the whole procedure’s held-out loss (it is built on the crossfit machinery; see$crossfit).
A first cross-fitted model
cf <- crossfit_super_learner(
data = mtcars,
formulas = mpg ~ hp + wt + am,
learners = list(lm = lnr_lm, mean = lnr_mean),
n_folds = 5
)
#> Warning: package 'future' was built under R version 4.5.2
cf
#> Cross-fitted Super Learner (nadir_crossfit_sl)
#> outcome: mpg (continuous)
#> outer folds: 5 inner CV folds: 5
#> observations: 32
#> cross-fitted loss on held-out data: 7.3264
#> Access: $oof_predict(data), $oof_predictions, $oof_predict_modified(modify), $oof_predict_fold(newdata_list),
#> $sl_fits, $training_data, $fold_assignments, $cv_lossThe honest predictions live in $oof_predictions, a plain
numeric vector in the row order of the input data:
head(cf$oof_predictions)
#> [1] 24.61570 24.11387 26.11689 20.48002 17.26625 20.41189
cf$cv_loss # empirical loss of the cross-fitted predictions
#> [1] 7.326419Because each outer fold estimates its own ensemble,
coef() returns a folds-by-learners weight matrix, and
summary() reports how stable those weights are across folds
— a useful diagnostic before trusting the cross-fitted predictions
downstream:
coef(cf)
#> lm mean
#> fold_1 0.9847008 0.01529921
#> fold_2 0.9197765 0.08022354
#> fold_3 0.9217920 0.07820799
#> fold_4 0.9271928 0.07280721
#> fold_5 0.9757505 0.02424947
summary(cf)
#> Summary of Cross-fitted Super Learner
#> outcome: mpg (continuous); n = 32
#> outer folds: 5; inner CV folds: 5
#>
#> Held-out loss by outer fold:
#> fold n_validation loss
#> 1 7 3.194
#> 2 6 9.007
#> 3 7 4.987
#> 4 6 7.068
#> 5 6 13.450
#> Overall cross-fitted loss: 7.326
#>
#> Ensemble weight stability across outer folds:
#> learner mean_weight sd_weight min_weight max_weight n_folds_present
#> lm 0.9458 0.0317 0.9198 0.9847 5
#> mean 0.0542 0.0317 0.0153 0.0802 5fitted(cf) and residuals(cf) are the
familiar S3 spellings of the same quantities: fitted() is
$oof_predictions, and residuals() is observed
minus out-of-fold predicted, both in input row order.
The oof_predict family
The same four-method interface is available on both
super_learner() fits (class nadir_sl_model)
and crossfit_super_learner() fits (class
nadir_crossfit_sl):
-
$oof_predictions— the stored vector of out-of-fold predictions, in input row order. Rows never held out by a non-coveringcv_schemaareNA. -
$oof_predict(newdata, rowids = NULL)— out-of-fold predictions for (possibly modified) versions of the training rows: row ofnewdatais predicted by the fold-specific fit that never saw row . With no arguments it returns the stored vector. -
$oof_predict_modified(modify)— the counterfactual workhorse: appliesmodify(newdata)to each held-out fold before predicting.$oof_predict_modified(NULL)re-predicts the unmodified data and agrees with$oof_predictions. -
$oof_predict_fold(newdata_list)— per-fold prediction vectors, for workflows that need the fold structure explicitly.
A note on row matching in $oof_predict(): out-of-fold
prediction is only meaningful for rows the model was cross-fit on, so
newdata rows must be matched back to training rows. If you
supplied rowids at fit time, matching is by id — shuffled
or subset newdata is fine, and rowids is then
required at predict time. Without fit-time rowids, rows are
matched by position and a one-time warning reminds you of the
assumption. We highly recommend using the rowids.
cf_ids <- crossfit_super_learner(
data = mtcars,
formulas = mpg ~ hp + wt + am,
learners = list(lm = lnr_lm, mean = lnr_mean),
n_folds = 5,
rowids = rownames(mtcars)
)
# a shuffled subset, matched by id for demonstration purposes
shuffled <- mtcars[c("Valiant", "Fiat 128", "Lotus Europa"), ]
cf_ids$oof_predict(shuffled, rowids = rownames(shuffled))
#> [1] 19.89846 25.71006 26.54238Honest nuisances for causal inference
The point of all this machinery is estimators built from cross-fitted
nuisance functions. Here is a complete AIPW estimate of an average
treatment effect: an outcome model
and a propensity model
,
each cross-fit, with the counterfactual predictions
and
obtained through $oof_predict_modified().
n <- 800
W1 <- rnorm(n)
W2 <- rnorm(n)
A <- rbinom(n, 1, plogis(0.4 * W1 - 0.3 * W2))
Y <- 1 + A + 0.8 * W1 - 0.5 * W2 + rnorm(n) # true ATE = 1
d <- data.frame(W1 = W1, W2 = W2, A = A, Y = Y)
cf_Q <- crossfit_super_learner(
data = d, formulas = Y ~ A + W1 + W2,
learners = list(lm = lnr_lm, mean = lnr_mean),
n_folds = 5
)
cf_g <- crossfit_super_learner(
data = d, formulas = A ~ W1 + W2,
learners = list(logistic = lnr_logistic),
outcome_type = "binary",
n_folds = 5
)
bound <- function(x, l = 0.005) pmin(pmax(x, l), 1 - l)
QAW <- cf_Q$oof_predictions # m(A_i, W_i)
Q1W <- cf_Q$oof_predict_modified(function(dd) {
dd$A <- 1
dd
}) # m(1, W_i)
Q0W <- cf_Q$oof_predict_modified(function(dd) {
dd$A <- 0
dd
}) # m(0, W_i)
gW <- bound(cf_g$oof_predictions) # g(W_i)
H <- A / gW - (1 - A) / (1 - gW)
phi <- (Q1W - Q0W) + H * (Y - QAW) # efficient influence function + psi
psi <- mean(phi)
se <- sd(phi) / sqrt(n)
c(estimate = psi, lower = psi - 1.96 * se, upper = psi + 1.96 * se)
#> estimate lower upper
#> 0.9042436 0.7625851 1.0459021Every quantity entering phi for observation
was produced by fits — learners and ensemble weights both — which were
estimated without observation
,
and that is what makes the nice efficient-influence-function based
standard error valid to use for inference.
Two refinements are provided for production use. First, estimators
like CV-TMLE want the outcome and propensity models cross-fit on the
same folds; pass both fits a shared custom
cv_schema that fixes one fold assignment (the package’s
test suite contains a complete shared-fold CV-TMLE recipe). Second, if
you are fitting super_learner() purely for its out-of-fold
interfaces, train_on_whole_dataset = FALSE skips the final
whole-data learner fits — saving roughly 1/n_folds of the
compute — while keeping $oof_predictions,
$oof_predict(), $oof_predict_modified(), and
$oof_predict_fold() fully functional and available for
use.