Cross-validation and hyperparameter search
tuning.Rdcross_validate() fits a model to each split's training rows and scores
it on the test rows. grid_search() does that for each candidate and
picks the lowest mean score; random_search() draws the candidates.
Models are built by your own fit function, so any of glm_fit(),
elastic_net_fit() or gam_fit() (or anything else) can be tuned.
Usage
cross_validate(data, splits, fit, score)
grid_search(candidates, data, splits, fit, score)
random_search(n, draw, data, splits, fit, score, seed = NULL)Arguments
- data
A data frame.
- splits
Splits from
k_fold(),group_k_fold()ortime_ordered().- fit
fit(train)forcross_validate(),fit(candidate, train)for the searches: returns a fitted model.- score
score(model, test): a loss on the test rows (lower is better), such as a meanfamily_deviance().- candidates
A list (or vector) of hyperparameter values.
- n
Number of random candidates.
- draw
draw(): one random candidate, using R's generator.- seed
Seed for R's generator before drawing.
Value
cross_validate(): one score per split. The searches: a data
frame with a candidate list column and the score, and attribute
best, the row of the lowest score.
Examples
d <- data.frame(x = 1:40 / 10)
d$y <- 1 + 2 * d$x + sin(1:40)
folds <- k_fold(nrow(d), 4, seed = 1)
mse <- function(m, test) mean((test$y - predict(m, test))^2)
cross_validate(d, folds, function(train) glm_fit(y ~ x, train, family = "gaussian"), mse)
#> [1] 0.6624610 0.7921088 0.4019159 0.5704141
g <- grid_search(c(0, 0.1, 1), d, folds,
function(lam, train) elastic_net_fit(y ~ x, train, family = "gaussian",
lambda = lam),
mse)
g[attr(g, "best"), ]
#> candidate score
#> 1 0 0.606725