43 double truncation_fraction) {
44 if (metrics.size() < 2) {
45 throw std::invalid_argument(
"ComputeTruncationGroups: metrics must have at least 2 entries");
47 if (!(truncation_fraction > 0.0) || truncation_fraction > 0.5) {
48 throw std::invalid_argument(
"ComputeTruncationGroups: truncation_fraction must be in (0, 0.5]");
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));
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]; });
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());
75 const std::map<std::string, double>& factors) {
78 auto config_it = config.find(spec.name);
79 if (config_it == config.end()) {
80 throw std::invalid_argument(
"ExploreConfigurationGivenFactors: config missing parameter '" +
86 auto factor_it = factors.find(spec.name);
87 if (factor_it == factors.end()) {
88 throw std::invalid_argument(
"ExploreConfigurationGivenFactors: factors missing parameter '" +
91 double factor = factor_it->second;
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));
97 double current = std::get<double>(config_it->second);
98 result[spec.name] = std::clamp(current * factor, spec.lower, spec.upper);
109template <
typename RNG>
111 std::bernoulli_distribution coin(0.5);
112 std::map<std::string, double> factors;
117 factors[spec.name] = coin(rng) ? 1.2 : 0.8;
132template <
typename RNG>
134 int num_epochs,
double truncation_fraction, RNG& rng) {
135 if (trials.empty()) {
136 throw std::invalid_argument(
"RunPBTGeneration: trials must not be empty");
138 if (num_epochs <= 0) {
139 throw std::invalid_argument(
"RunPBTGeneration: num_epochs must be positive");
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);
147 if (trials.size() < 2) {
152 std::uniform_int_distribution<size_t> pick_top(0, groups.
top_indices.size() - 1);
154 size_t source_index = groups.
top_indices[pick_top(rng)];
155 if (source_index == bottom_index) {
158 trials[bottom_index]->SetWeights(trials[source_index]->GetWeights());
159 Configuration copied = trials[source_index]->GetHyperparameters();
161 metrics[bottom_index] = metrics[source_index];
179template <
typename RNG>
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");
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);
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};
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