pulsatrix
Loading...
Searching...
No Matches
asha.hpp
Go to the documentation of this file.
1
18#pragma once
19
20#include <algorithm>
21#include <functional>
22#include <limits>
23#include <memory>
24#include <stdexcept>
25#include <vector>
26
28
29namespace pulsatrix {
30
39
40namespace detail {
43 std::unique_ptr<ResumableTrial> trial;
44 int rung;
45 double metric;
46};
47} // namespace detail
48
67inline ASHAResult RunASHAOnConfigQueue(std::vector<Configuration> config_queue,
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");
72 }
73 if (initial_epoch_budget <= 0) {
74 throw std::invalid_argument("RunASHAOnConfigQueue: initial_epoch_budget must be positive");
75 }
76 if (eta <= 1.0) {
77 throw std::invalid_argument("RunASHAOnConfigQueue: eta must be > 1.0");
78 }
79 if (num_rungs < 2) {
80 throw std::invalid_argument("RunASHAOnConfigQueue: num_rungs must be >= 2");
81 }
82
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);
87 }
88
89 std::vector<std::vector<double>> rung_history(static_cast<size_t>(num_rungs));
90 std::vector<detail::ASHACandidate> candidates;
91 size_t queue_pos = 0;
92 size_t total_epochs_trained = 0;
93
94 while (true) {
95 int promote_rung = -1;
96 int promote_idx = -1;
97
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) {
101 continue;
102 }
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];
107
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;
115 }
116 }
117 }
118 if (local_best_idx != -1) {
119 promote_rung = k;
120 promote_idx = local_best_idx;
121 break; // highest-k promotable candidate wins -- prefer finishing over starting fresh
122 }
123 }
124
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);
133 continue;
134 }
135
136 if (queue_pos < config_queue.size()) {
137 Configuration config = std::move(config_queue[queue_pos++]);
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);
142 candidates.push_back(detail::ASHACandidate{std::move(config), std::move(trial), 0, metric});
143 continue;
144 }
145
146 break; // nothing left to promote or start
147 }
148
149 auto best_it = std::max_element(
150 candidates.begin(), candidates.end(),
151 [](const detail::ASHACandidate& a, const detail::ASHACandidate& b) { return a.metric < b.metric; });
152 return ASHAResult{best_it->config, best_it->metric, total_epochs_trained, queue_pos};
153}
154
160template <typename RNG>
161ASHAResult RunASHA(const SearchSpace& space, const TrialFactory& make_trial, size_t max_configs_started,
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");
165 }
166 std::vector<Configuration> queue;
167 queue.reserve(max_configs_started);
168 for (size_t i = 0; i < max_configs_started; ++i) {
169 queue.push_back(RandomSample(space, rng));
170 }
171 return RunASHAOnConfigQueue(std::move(queue), make_trial, initial_epoch_budget, eta, num_rungs);
172}
173
174} // 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
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
The best configuration/metric found, the total epoch budget spent, and how many distinct configuratio...
Definition asha.hpp:33
size_t num_configs_started
Definition asha.hpp:37
size_t total_epochs_trained
Definition asha.hpp:36
Configuration best_configuration
Definition asha.hpp:34
double best_metric
Definition asha.hpp:35
Definition asha.hpp:41
double metric
Definition asha.hpp:45
int rung
Definition asha.hpp:44
std::unique_ptr< ResumableTrial > trial
Definition asha.hpp:43
Configuration config
Definition asha.hpp:42
Successive Halving (the rung-based promotion mechanism underlying Hyperband and ASHA,...