pulsatrix
Loading...
Searching...
No Matches
pbt.hpp
Go to the documentation of this file.
1
11#pragma once
12
13#include <algorithm>
14#include <cmath>
15#include <map>
16#include <memory>
17#include <numeric>
18#include <random>
19#include <stdexcept>
20#include <variant>
21#include <vector>
22
25
26namespace pulsatrix {
27
30 std::vector<size_t> bottom_indices;
31 std::vector<size_t> top_indices;
32};
33
42inline PBTTruncationGroups ComputeTruncationGroups(const std::vector<double>& metrics,
43 double truncation_fraction) {
44 if (metrics.size() < 2) {
45 throw std::invalid_argument("ComputeTruncationGroups: metrics must have at least 2 entries");
46 }
47 if (!(truncation_fraction > 0.0) || truncation_fraction > 0.5) {
48 throw std::invalid_argument("ComputeTruncationGroups: truncation_fraction must be in (0, 0.5]");
49 }
50
51 size_t n = metrics.size();
52 size_t num_selected = std::max<size_t>(1, static_cast<size_t>(static_cast<double>(n) * truncation_fraction));
53
54 std::vector<size_t> order(n);
55 std::iota(order.begin(), order.end(), 0);
56 std::stable_sort(order.begin(), order.end(), [&](size_t a, size_t b) { return metrics[a] > metrics[b]; });
57
59 groups.top_indices.assign(order.begin(), order.begin() + static_cast<long>(num_selected));
60 groups.bottom_indices.assign(order.end() - static_cast<long>(num_selected), order.end());
61 return groups;
62}
63
75 const std::map<std::string, double>& factors) {
76 Configuration result = config;
77 for (const auto& spec : space.parameters()) {
78 auto config_it = config.find(spec.name);
79 if (config_it == config.end()) {
80 throw std::invalid_argument("ExploreConfigurationGivenFactors: config missing parameter '" +
81 spec.name + "'");
82 }
83 if (spec.kind == ParameterKind::Categorical) {
84 continue;
85 }
86 auto factor_it = factors.find(spec.name);
87 if (factor_it == factors.end()) {
88 throw std::invalid_argument("ExploreConfigurationGivenFactors: factors missing parameter '" +
89 spec.name + "'");
90 }
91 double factor = factor_it->second;
92 if (spec.kind == ParameterKind::Integer) {
93 int64_t current = std::get<int64_t>(config_it->second);
94 double perturbed = std::clamp(static_cast<double>(current) * factor, spec.lower, spec.upper);
95 result[spec.name] = static_cast<int64_t>(std::llround(perturbed));
96 } else {
97 double current = std::get<double>(config_it->second);
98 result[spec.name] = std::clamp(current * factor, spec.lower, spec.upper);
99 }
100 }
101 return result;
102}
103
109template <typename RNG>
110Configuration ExploreConfiguration(const Configuration& config, const SearchSpace& space, RNG& rng) {
111 std::bernoulli_distribution coin(0.5);
112 std::map<std::string, double> factors;
113 for (const auto& spec : space.parameters()) {
114 if (spec.kind == ParameterKind::Categorical) {
115 continue;
116 }
117 factors[spec.name] = coin(rng) ? 1.2 : 0.8;
118 }
119 return ExploreConfigurationGivenFactors(config, space, factors);
120}
121
132template <typename RNG>
133std::vector<double> RunPBTGeneration(std::vector<std::unique_ptr<PBTResumableTrial>>& trials, const SearchSpace& space,
134 int num_epochs, double truncation_fraction, RNG& rng) {
135 if (trials.empty()) {
136 throw std::invalid_argument("RunPBTGeneration: trials must not be empty");
137 }
138 if (num_epochs <= 0) {
139 throw std::invalid_argument("RunPBTGeneration: num_epochs must be positive");
140 }
141
142 std::vector<double> metrics(trials.size());
143 for (size_t i = 0; i < trials.size(); ++i) {
144 metrics[i] = trials[i]->TrainForEpochs(num_epochs);
145 }
146
147 if (trials.size() < 2) {
148 return metrics; // nothing to exploit/explore against with a single trial
149 }
150
151 PBTTruncationGroups groups = ComputeTruncationGroups(metrics, truncation_fraction);
152 std::uniform_int_distribution<size_t> pick_top(0, groups.top_indices.size() - 1);
153 for (size_t bottom_index : groups.bottom_indices) {
154 size_t source_index = groups.top_indices[pick_top(rng)];
155 if (source_index == bottom_index) {
156 continue; // only possible if top/bottom overlap at a tiny population size
157 }
158 trials[bottom_index]->SetWeights(trials[source_index]->GetWeights());
159 Configuration copied = trials[source_index]->GetHyperparameters();
160 trials[bottom_index]->SetHyperparameters(ExploreConfiguration(copied, space, rng));
161 metrics[bottom_index] = metrics[source_index];
162 }
163
164 return metrics;
165}
166
173
179template <typename RNG>
180PBTResult RunPBT(std::vector<std::unique_ptr<PBTResumableTrial>>& trials, const SearchSpace& space,
181 int num_generations, int epochs_per_generation, double truncation_fraction, RNG& rng) {
182 if (num_generations <= 0) {
183 throw std::invalid_argument("RunPBT: num_generations must be positive");
184 }
185
186 std::vector<double> metrics;
187 for (int gen = 0; gen < num_generations; ++gen) {
188 metrics = RunPBTGeneration(trials, space, epochs_per_generation, truncation_fraction, rng);
189 }
190
191 size_t best_index = static_cast<size_t>(std::max_element(metrics.begin(), metrics.end()) - metrics.begin());
192 return PBTResult{best_index, metrics[best_index], num_generations};
193}
194
195} // namespace pulsatrix
Describes a hyperparameter search space as an ordered list of named, typed parameters....
Definition search_space.hpp:54
const std::vector< ParameterSpec > & parameters() const
Every parameter, in the order added.
Definition search_space.hpp:116
Definition acquisition_functions.hpp:16
PBTResult RunPBT(std::vector< std::unique_ptr< PBTResumableTrial > > &trials, const SearchSpace &space, int num_generations, int epochs_per_generation, double truncation_fraction, RNG &rng)
Runs num_generations of RunPBTGeneration in sequence.
Definition pbt.hpp:180
Configuration ExploreConfiguration(const Configuration &config, const SearchSpace &space, RNG &rng)
RNG-driven wrapper: draws each non-categorical parameter's factor uniformly from {0....
Definition pbt.hpp:110
Configuration ExploreConfigurationGivenFactors(const Configuration &config, const SearchSpace &space, const std::map< std::string, double > &factors)
Pure core: applies an explicit per-parameter multiplicative factor to every Continuous/LogUniform/Int...
Definition pbt.hpp:74
PBTTruncationGroups ComputeTruncationGroups(const std::vector< double > &metrics, double truncation_fraction)
Pure core: identifies the bottom and top truncation_fraction of the population by metric (higher is b...
Definition pbt.hpp:42
std::map< std::string, ConfigValue > Configuration
A concrete hyperparameter configuration: parameter name -> concrete value.
Definition search_space.hpp:46
std::vector< double > RunPBTGeneration(std::vector< std::unique_ptr< PBTResumableTrial > > &trials, const SearchSpace &space, int num_epochs, double truncation_fraction, RNG &rng)
Runs one PBT generation: trains every live trial for num_epochs, then exploits+explores the bottom tr...
Definition pbt.hpp:133
Extends the sibling HPO campaign's own ResumableTrial (successive_halving.hpp) with the weight/hyperp...
Typed hyperparameter search-space description: named parameters, each continuous, log-uniform,...
The best-performing trial's index, its metric, and how many generations ran.
Definition pbt.hpp:168
size_t best_trial_index
Definition pbt.hpp:169
double best_metric
Definition pbt.hpp:170
int generations_run
Definition pbt.hpp:171
Indices of the population's current worst- and best-performing members.
Definition pbt.hpp:29
std::vector< size_t > bottom_indices
Definition pbt.hpp:30
std::vector< size_t > top_indices
Definition pbt.hpp:31