Skip to content

[source]

Grid Search#

Grid Search is an algorithm that optimizes hyper-parameter selection. From the user's perspective, the process of training and predicting is the same, however, under the hood Grid Search trains a model for each combination of possible parameters and the best model is selected as the base estimator.

Interfaces: Estimator, Learner, Parallel, Persistable, Verbose

Data Type Compatibility: Depends on base learner

Parameters#

# Name Default Type Description
1 class string The class name of the base learner.
2 params array An array of lists containing the possible values for each of the base learner's constructor parameters.
3 metric auto Metric The validation metric used to score each set of hyper-parameters.
4 validator KFold Validator The validator used to test and score the model.

Example#

use Rubix\ML\GridSearch;
use Rubix\ML\Classifiers\KNearestNeighbors;
use Rubix\ML\Kernels\Distance\Euclidean;
use Rubix\ML\Kernels\Distance\Manhattan;
use Rubix\ML\CrossValidation\Metrics\FBeta;
use Rubix\ML\CrossValidation\KFold;

$params = [
    [1, 3, 5, 10], [true, false], [new Euclidean(), new Manhattan()],
];

$estimator = new GridSearch(KNearestNeighbors::class, $params, new FBeta(), new KFold(5));

Passing an empty array [] for any of the base learner's constructor parameters tells Grid Search to use that parameter's default value from the base learner's constructor (or null if no default exists).

You can also construct a Grid Search instance via the fromNamedParams() factory. Specify the hyper-parameters by the name of the base learner's constructor parameter (order does not matter). Hyper-parameters that are omitted are assigned their default value from the base learner's constructor.

$estimator = GridSearch::fromNamedParams(
    KNearestNeighbors::class,
    [
        'kernel' => [new Euclidean(), new Manhattan()],
        'k' => [1, 3, 5, 10],
        'weighted' => [true, false],
    ],
    new FBeta(),
    new KFold(5)
);

Setup#

Some estimators expose configuration methods that are unrelated to their constructor parameters—for example, setValidationDataset() on iterative learners. Use setup() to register a callback that is invoked on each newly-instantiated base estimator before it is cross-validated (and again on the final best estimator before it is trained on the full dataset).

The callback receives the base estimator instance and may call any of its methods. It is not serialized with the grid search instance.

$estimator = new GridSearch(LogisticRegression::class, $params, new FBeta(), new KFold(5));

$estimator->setup(function (LogisticRegression $regressor) : void {
    $regressor->setValidationDataset($validation);
});

Parallel#

This estimator implements the Parallel interface and can utilize a parallel processing backend such as Amp to speed up training and inference:

use Rubix\ML\Backends\Amp;

$estimator->setBackend(new Amp(4));

Additional Methods#

Return the base learner instance.

public base() : \Rubix\ML\Estimator

Return a Report of every parameter combination tested along with its validation score from the last search. The rows are keyed by trial number in the order the trials were trained in, so the Trial N entries line up with the N-th parameter combination (scores()[N - 1]) and the trial numbers logged during training. Each row pairs the validation score (keyed by the metric name) with a nested params map of the combination's constructor parameters:

public results() : \Rubix\ML\Report

Return the validation scores of each of the parameter combinations. The scores are returned in the order the trials were trained in, so the score at index N - 1 corresponds to Trial N.

public scores() : ?array

Return all the parameter combinations.

public combinations() : array