86 int initial_epoch_budget,
double eta) {
87 if (configs.empty()) {
88 throw std::invalid_argument(
"RunSuccessiveHalvingOnConfigs: configs must not be empty");
90 if (initial_epoch_budget <= 0) {
91 throw std::invalid_argument(
"RunSuccessiveHalvingOnConfigs: initial_epoch_budget must be positive");
94 throw std::invalid_argument(
"RunSuccessiveHalvingOnConfigs: eta must be > 1.0");
99 std::unique_ptr<ResumableTrial> trial;
102 std::vector<Candidate> candidates;
103 candidates.reserve(configs.size());
104 for (
auto& config : configs) {
105 auto trial = make_trial(config);
106 candidates.push_back(Candidate{std::move(config), std::move(trial),
107 -std::numeric_limits<double>::infinity()});
110 size_t total_epochs_trained = 0;
111 int budget = initial_epoch_budget;
112 bool trained_at_least_once =
false;
115 for (
auto& c : candidates) {
116 c.metric = c.trial->TrainForEpochs(budget);
118 total_epochs_trained += candidates.size() *
static_cast<size_t>(budget);
119 trained_at_least_once =
true;
121 if (candidates.size() <= 1) {
124 std::sort(candidates.begin(), candidates.end(),
125 [](
const Candidate& a,
const Candidate& b) { return a.metric > b.metric; });
127 size_t num_survivors = std::max<size_t>(1,
static_cast<size_t>(candidates.size() / eta));
128 if (num_survivors >= candidates.size()) {
131 candidates.resize(num_survivors);
132 budget =
static_cast<int>(
static_cast<double>(budget) * eta);
134 (void)trained_at_least_once;
136 auto best_it = std::max_element(candidates.begin(), candidates.end(),
137 [](
const Candidate& a,
const Candidate& b) { return a.metric < b.metric; });
147template <
typename RNG>
149 size_t num_configs,
int initial_epoch_budget,
double eta,
151 if (num_configs == 0) {
152 throw std::invalid_argument(
"RunSuccessiveHalving: num_configs must be positive");
154 std::vector<Configuration> configs;
155 configs.reserve(num_configs);
156 for (
size_t i = 0; i < num_configs; ++i) {
A single hyperparameter configuration's live, resumable training state – own whatever network/optimiz...
Definition successive_halving.hpp:50
virtual double TrainForEpochs(int num_epochs)=0
Trains this trial for num_epochs additional epochs (continuing from wherever this trial's own trainin...
virtual ~ResumableTrial()=default
Describes a hyperparameter search space as an ordered list of named, typed parameters....
Definition search_space.hpp:54
Grid search and random search over a SearchSpace: the two simplest, baseline hyperparameter-optimizat...
Definition acquisition_functions.hpp:16
std::map< std::string, ConfigValue > Configuration
A concrete hyperparameter configuration: parameter name -> concrete value.
Definition search_space.hpp:46
std::function< std::unique_ptr< ResumableTrial >(const Configuration &)> TrialFactory
Builds a fresh ResumableTrial for a given configuration.
Definition successive_halving.hpp:64
SuccessiveHalvingResult RunSuccessiveHalvingOnConfigs(std::vector< Configuration > configs, const TrialFactory &make_trial, int initial_epoch_budget, double eta)
Runs Successive Halving over an explicit, caller-supplied list of configurations – the pure,...
Definition successive_halving.hpp:84
Configuration RandomSample(const SearchSpace &space, RNG &rng)
Draws one configuration uniformly at random from space: Continuous parameters uniform over [lower,...
Definition hpo_sampling.hpp:27
SuccessiveHalvingResult RunSuccessiveHalving(const SearchSpace &space, const TrialFactory &make_trial, size_t num_configs, int initial_epoch_budget, double eta, RNG &rng)
RNG-driven wrapper: draws num_configs configurations from space via RandomSample, then runs RunSucces...
Definition successive_halving.hpp:148
Typed hyperparameter search-space description: named parameters, each continuous, log-uniform,...
The winning configuration, its final metric, and the total epoch-budget actually spent across every t...
Definition successive_halving.hpp:69
Configuration best_configuration
Definition successive_halving.hpp:70
double best_metric
Definition successive_halving.hpp:71
size_t total_epochs_trained
Definition successive_halving.hpp:72