68 const TrialFactory& make_trial,
int initial_epoch_budget,
69 double eta,
int num_rungs) {
70 if (config_queue.empty()) {
71 throw std::invalid_argument(
"RunASHAOnConfigQueue: config_queue must not be empty");
73 if (initial_epoch_budget <= 0) {
74 throw std::invalid_argument(
"RunASHAOnConfigQueue: initial_epoch_budget must be positive");
77 throw std::invalid_argument(
"RunASHAOnConfigQueue: eta must be > 1.0");
80 throw std::invalid_argument(
"RunASHAOnConfigQueue: num_rungs must be >= 2");
83 std::vector<int> budgets(
static_cast<size_t>(num_rungs));
84 budgets[0] = initial_epoch_budget;
85 for (
int i = 1; i < num_rungs; ++i) {
86 budgets[
static_cast<size_t>(i)] =
static_cast<int>(budgets[
static_cast<size_t>(i - 1)] * eta);
89 std::vector<std::vector<double>> rung_history(
static_cast<size_t>(num_rungs));
90 std::vector<detail::ASHACandidate> candidates;
92 size_t total_epochs_trained = 0;
95 int promote_rung = -1;
98 for (
int k = num_rungs - 2; k >= 0; --k) {
99 auto& history = rung_history[
static_cast<size_t>(k)];
100 if (
static_cast<double>(history.size()) < eta) {
103 std::vector<double> sorted_desc = history;
104 std::sort(sorted_desc.begin(), sorted_desc.end(), std::greater<double>());
105 size_t keep = std::max<size_t>(1,
static_cast<size_t>(
static_cast<double>(sorted_desc.size()) / eta));
106 double cutoff = sorted_desc[keep - 1];
108 int local_best_idx = -1;
109 double local_best_metric = -std::numeric_limits<double>::infinity();
110 for (
size_t i = 0; i < candidates.size(); ++i) {
111 if (candidates[i].rung == k && candidates[i].metric >= cutoff) {
112 if (local_best_idx == -1 || candidates[i].metric > local_best_metric) {
113 local_best_idx =
static_cast<int>(i);
114 local_best_metric = candidates[i].metric;
118 if (local_best_idx != -1) {
120 promote_idx = local_best_idx;
125 if (promote_idx != -1) {
126 int k = promote_rung;
127 int additional = budgets[
static_cast<size_t>(k + 1)] - budgets[
static_cast<size_t>(k)];
128 double new_metric = candidates[
static_cast<size_t>(promote_idx)].trial->TrainForEpochs(additional);
129 total_epochs_trained +=
static_cast<size_t>(additional);
130 candidates[
static_cast<size_t>(promote_idx)].metric = new_metric;
131 candidates[
static_cast<size_t>(promote_idx)].rung = k + 1;
132 rung_history[
static_cast<size_t>(k + 1)].push_back(new_metric);
136 if (queue_pos < config_queue.size()) {
138 auto trial = make_trial(config);
139 double metric = trial->TrainForEpochs(budgets[0]);
140 total_epochs_trained +=
static_cast<size_t>(budgets[0]);
141 rung_history[0].push_back(metric);
149 auto best_it = std::max_element(
150 candidates.begin(), candidates.end(),
152 return ASHAResult{best_it->config, best_it->metric, total_epochs_trained, queue_pos};
162 int initial_epoch_budget,
double eta,
int num_rungs, RNG& rng) {
163 if (max_configs_started == 0) {
164 throw std::invalid_argument(
"RunASHA: max_configs_started must be positive");
166 std::vector<Configuration> queue;
167 queue.reserve(max_configs_started);
168 for (
size_t i = 0; i < max_configs_started; ++i) {
171 return RunASHAOnConfigQueue(std::move(queue), make_trial, initial_epoch_budget, eta, num_rungs);
ASHAResult RunASHAOnConfigQueue(std::vector< Configuration > config_queue, const TrialFactory &make_trial, int initial_epoch_budget, double eta, int num_rungs)
Runs ASHA, drawing new configurations from an explicit, ordered queue (rather than sampling indefinit...
Definition asha.hpp:67
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
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
ASHAResult RunASHA(const SearchSpace &space, const TrialFactory &make_trial, size_t max_configs_started, int initial_epoch_budget, double eta, int num_rungs, RNG &rng)
RNG-driven wrapper: draws max_configs_started configurations from space via RandomSample to serve as ...
Definition asha.hpp:161