pulsatrix
Loading...
Searching...
No Matches
hyperband.hpp
Go to the documentation of this file.
1
20#pragma once
21
22#include <cmath>
23#include <limits>
24#include <stdexcept>
25#include <vector>
26
28
29namespace pulsatrix {
30
35 int s;
38};
39
47inline std::vector<HyperbandBracket> ComputeHyperbandBrackets(int max_resource, double eta) {
48 if (max_resource <= 0) {
49 throw std::invalid_argument("ComputeHyperbandBrackets: max_resource must be positive");
50 }
51 if (eta <= 1.0) {
52 throw std::invalid_argument("ComputeHyperbandBrackets: eta must be > 1.0");
53 }
54
55 int s_max = static_cast<int>(std::floor(std::log(static_cast<double>(max_resource)) / std::log(eta)));
56 double b = static_cast<double>(s_max + 1) * static_cast<double>(max_resource);
57
58 std::vector<HyperbandBracket> brackets;
59 for (int s = s_max; s >= 0; --s) {
60 double eta_pow_s = std::pow(eta, s);
61 size_t num_configs =
62 std::max<size_t>(1, static_cast<size_t>(std::ceil((b / max_resource) * (eta_pow_s / (s + 1)))));
63 int initial_budget = std::max(1, static_cast<int>(std::round(max_resource / eta_pow_s)));
64 brackets.push_back(HyperbandBracket{s, num_configs, initial_budget});
65 }
66 return brackets;
67}
68
76
83template <typename RNG>
84HyperbandResult RunHyperband(const SearchSpace& space, const TrialFactory& make_trial, int max_resource,
85 double eta, RNG& rng) {
86 auto brackets = ComputeHyperbandBrackets(max_resource, eta);
87
88 HyperbandResult overall{Configuration{}, -std::numeric_limits<double>::infinity(), 0};
89 for (const auto& bracket : brackets) {
90 auto result = RunSuccessiveHalving(space, make_trial, bracket.num_configs, bracket.initial_budget, eta, rng);
91 overall.total_epochs_trained += result.total_epochs_trained;
92 if (result.best_metric > overall.best_metric) {
93 overall.best_metric = result.best_metric;
94 overall.best_configuration = result.best_configuration;
95 }
96 }
97 return overall;
98}
99
100} // namespace pulsatrix
Describes a hyperparameter search space as an ordered list of named, typed parameters....
Definition search_space.hpp:54
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::vector< HyperbandBracket > ComputeHyperbandBrackets(int max_resource, double eta)
Computes the classic Hyperband bracket schedule: s_max = floor(log_eta(max_resource)),...
Definition hyperband.hpp:47
std::function< std::unique_ptr< ResumableTrial >(const Configuration &)> TrialFactory
Builds a fresh ResumableTrial for a given configuration.
Definition successive_halving.hpp:64
HyperbandResult RunHyperband(const SearchSpace &space, const TrialFactory &make_trial, int max_resource, double eta, RNG &rng)
Runs one Successive Halving bracket (successive_halving.hpp) per ComputeHyperbandBrackets(max_resourc...
Definition hyperband.hpp:84
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
One bracket's own (num_configs, initial_budget) trade-off point; s is the bracket index (s_max = most...
Definition hyperband.hpp:34
size_t num_configs
Definition hyperband.hpp:36
int initial_budget
Definition hyperband.hpp:37
int s
Definition hyperband.hpp:35
The best configuration/metric found across every bracket, and the total epoch budget spent summed acr...
Definition hyperband.hpp:71
double best_metric
Definition hyperband.hpp:73
Configuration best_configuration
Definition hyperband.hpp:72
size_t total_epochs_trained
Definition hyperband.hpp:74
Successive Halving (the rung-based promotion mechanism underlying Hyperband and ASHA,...