ppforest2 v0.1.3
Projection Pursuit Decision Trees and Random Forests
Loading...
Searching...
No Matches
TrainingSpec.hpp
Go to the documentation of this file.
1#pragma once
2
10
11#include <algorithm>
12#include <memory>
13#include <thread>
14#include <nlohmann/json.hpp>
15
16namespace ppforest2 {
45 public:
46 using Ptr = std::shared_ptr<TrainingSpec>;
47
62
65
67 int const size;
69 int const seed;
71 int const threads;
73 int const max_retries;
74
103 class Builder {
104 public:
126
129
131 : mode(mode) {}
132
134 config.pp = std::move(v);
135 return *this;
136 }
138 config.vars = std::move(v);
139 return *this;
140 }
142 config.cutpoint = std::move(v);
143 return *this;
144 }
146 config.stop = std::move(v);
147 return *this;
148 }
150 config.binarization = std::move(v);
151 return *this;
152 }
154 config.grouping = std::move(v);
155 return *this;
156 }
158 config.leaf = std::move(v);
159 return *this;
160 }
161
162 Builder& size(int v) {
163 config.size = v;
164 return *this;
165 }
166 Builder& seed(int v) {
167 config.seed = v;
168 return *this;
169 }
170 Builder& threads(int v) {
171 config.threads = v;
172 return *this;
173 }
175 config.max_retries = v;
176 return *this;
177 }
178
208
216
219 };
220
229
254 int size,
255 int seed,
256 int threads,
257 int max_retries
258 );
259
260 // -- Forwarding methods (delegate to the underlying strategy) -----------
261
263 void find_projection(NodeContext& ctx, stats::RNG& rng) const;
264
266 void select_vars(NodeContext& ctx, stats::RNG& rng) const;
267
269 void find_cutpoint(NodeContext& ctx, stats::RNG& rng) const;
270
278 bool should_stop(NodeContext const& ctx, stats::RNG& rng) const;
279
281 void regroup(NodeContext& ctx, stats::RNG& rng) const;
282
290 void group(NodeContext& ctx, stats::RNG& rng) const;
291
294
296 TreeNode::Ptr create_leaf(NodeContext const& ctx, stats::RNG& rng) const { return leaf->create_leaf(ctx, rng); }
297
299 bool is_forest() const { return size > 0; }
300
302 nlohmann::json to_json() const;
303
305 static Ptr from_json(nlohmann::json const& j);
306
308 template<typename... Args> static Ptr make(Args&&... args) {
309 return std::make_shared<TrainingSpec>(std::forward<Args>(args)...);
310 }
311
320 int resolve_threads() const {
321 // hardware_concurrency() may return 0 when it cannot be determined;
322 // passing a non-positive count to omp_set_num_threads is undefined.
323 return threads > 0 ? threads : std::max(1, static_cast<int>(std::thread::hardware_concurrency()));
324 }
325 };
326
328 inline bool is_classification(TrainingSpec const& spec) {
329 return types::is_classification(spec.mode);
330 }
331
333 inline bool is_regression(TrainingSpec const& spec) {
334 return types::is_regression(spec.mode);
335 }
336
338 inline bool is_classification(TrainingSpec::Ptr const& spec) {
339 return spec != nullptr && is_classification(*spec);
340 }
341
343 inline bool is_regression(TrainingSpec::Ptr const& spec) {
344 return spec != nullptr && is_regression(*spec);
345 }
346}
std::shared_ptr< ProjectionPursuit > Ptr
Definition Strategy.hpp:95
Fluent builder for TrainingSpec.
Definition TrainingSpec.hpp:103
Builder & leaf(leaf::LeafStrategy::Ptr v)
Definition TrainingSpec.hpp:157
Builder & seed(int v)
Definition TrainingSpec.hpp:166
Builder & cutpoint(cutpoint::Cutpoint::Ptr v)
Definition TrainingSpec.hpp:141
Builder & threads(int v)
Definition TrainingSpec.hpp:170
Builder & binarization(binarize::Binarization::Ptr v)
Definition TrainingSpec.hpp:149
Builder & max_retries(int v)
Definition TrainingSpec.hpp:174
Builder & size(int v)
Definition TrainingSpec.hpp:162
Ptr make()
Shorthand for std::make_shared<TrainingSpec>(build()).
Builder(types::Mode mode)
Definition TrainingSpec.hpp:130
types::Mode const mode
Definition TrainingSpec.hpp:128
Builder & pp(pp::ProjectionPursuit::Ptr v)
Definition TrainingSpec.hpp:133
TrainingSpec build()
Finalize the builder into a TrainingSpec.
Builder & vars(vars::VariableSelection::Ptr v)
Definition TrainingSpec.hpp:137
Builder & apply_defaults()
Fill in any null strategy fields with mode-aware defaults.
Config config
Definition TrainingSpec.hpp:127
Builder & stop(stop::StopRule::Ptr v)
Definition TrainingSpec.hpp:145
Builder & grouping(grouping::Grouping::Ptr v)
Definition TrainingSpec.hpp:153
Training configuration for projection pursuit trees and forests.
Definition TrainingSpec.hpp:44
TreeNode::Ptr create_leaf(NodeContext const &ctx, stats::RNG &rng) const
Create a leaf node from the current node context.
Definition TrainingSpec.hpp:296
grouping::Grouping::Ptr const grouping
Grouping strategy.
Definition TrainingSpec.hpp:59
static Builder builder(types::Mode mode)
Create a builder for the given mode.
Definition TrainingSpec.hpp:228
static Ptr make(Args &&... args)
Create a shared pointer to a TrainingSpec.
Definition TrainingSpec.hpp:308
TrainingSpec(pp::ProjectionPursuit::Ptr pp, vars::VariableSelection::Ptr vars, cutpoint::Cutpoint::Ptr cutpoint, stop::StopRule::Ptr stop, binarize::Binarization::Ptr binarization, grouping::Grouping::Ptr grouping, leaf::LeafStrategy::Ptr leaf, types::Mode mode, int size, int seed, int threads, int max_retries)
Construct a training specification.
void group(NodeContext &ctx, stats::RNG &rng) const
Split observations into two child partitions.
binarize::Binarization::Ptr const binarization
Binarization strategy.
Definition TrainingSpec.hpp:57
leaf::LeafStrategy::Ptr const leaf
Leaf creation strategy.
Definition TrainingSpec.hpp:61
pp::ProjectionPursuit::Ptr const pp
Projection pursuit optimization strategy.
Definition TrainingSpec.hpp:49
int resolve_threads() const
Get the number of threads to use for training.
Definition TrainingSpec.hpp:320
void find_projection(NodeContext &ctx, stats::RNG &rng) const
Run projection pursuit optimization. Asserts postcondition: ctx.projector and ctx....
nlohmann::json to_json() const
Serialize the training spec to JSON.
void find_cutpoint(NodeContext &ctx, stats::RNG &rng) const
Compute the split cutpoint. Asserts postcondition: ctx.cutpoint is set.
bool is_forest() const
Whether this specification describes a forest (size > 0).
Definition TrainingSpec.hpp:299
cutpoint::Cutpoint::Ptr const cutpoint
Split cutpoint strategy.
Definition TrainingSpec.hpp:53
stop::StopRule::Ptr const stop
Stop rule strategy.
Definition TrainingSpec.hpp:55
bool should_stop(NodeContext const &ctx, stats::RNG &rng) const
Check whether the node should stop growing.
vars::VariableSelection::Ptr const vars
Variable selection strategy.
Definition TrainingSpec.hpp:51
int const max_retries
Maximum retry attempts for degenerate trees.
Definition TrainingSpec.hpp:73
int const size
Number of trees (0 = single tree).
Definition TrainingSpec.hpp:67
void regroup(NodeContext &ctx, stats::RNG &rng) const
Reduce multiclass partition to binary. Asserts postcondition: ctx.y_bin is set.
int const seed
RNG seed.
Definition TrainingSpec.hpp:69
void select_vars(NodeContext &ctx, stats::RNG &rng) const
Run variable selection. Asserts postcondition: ctx.var_selection is set.
static Ptr from_json(nlohmann::json const &j)
Deserialize a training spec from JSON.
types::Mode const mode
Training mode (classification or regression).
Definition TrainingSpec.hpp:64
int const threads
Number of threads for parallel forest training.
Definition TrainingSpec.hpp:71
stats::GroupPartition init_groups(types::OutcomeVector const &y) const
Create the initial group partition from the training response.
Definition TrainingSpec.hpp:293
std::shared_ptr< TrainingSpec > Ptr
Definition TrainingSpec.hpp:46
std::unique_ptr< TreeNode > Ptr
Definition TreeNode.hpp:21
Contiguous-block representation of grouped observations.
Definition GroupPartition.hpp:40
pcg32 RNG
Definition Stats.hpp:24
bool is_classification(Mode mode)
Whether mode is Classification.
Definition Types.hpp:71
Eigen::Matrix< Outcome, Eigen::Dynamic, 1 > OutcomeVector
Dynamic-size column vector of predictions.
Definition Types.hpp:52
bool is_regression(Mode mode)
Whether mode is Regression.
Definition Types.hpp:76
Mode
Training mode.
Definition Types.hpp:68
Binarization strategies for multiclass-to-binary reduction.
Definition Benchmark.hpp:25
bool is_classification(Model const &model)
Whether model was trained for classification.
Definition Model.hpp:145
bool is_regression(Model const &model)
Whether model was trained for regression.
Definition Model.hpp:155
Mutable context accumulating intermediate results during node training.
Definition NodeContext.hpp:20
Builder state — the configuration being assembled.
Definition TrainingSpec.hpp:112
leaf::LeafStrategy::Ptr leaf
Definition TrainingSpec.hpp:119
int max_retries
Definition TrainingSpec.hpp:124
int size
Definition TrainingSpec.hpp:121
cutpoint::Cutpoint::Ptr cutpoint
Definition TrainingSpec.hpp:115
int seed
Definition TrainingSpec.hpp:122
grouping::Grouping::Ptr grouping
Definition TrainingSpec.hpp:118
int threads
Definition TrainingSpec.hpp:123
vars::VariableSelection::Ptr vars
Definition TrainingSpec.hpp:114
pp::ProjectionPursuit::Ptr pp
Definition TrainingSpec.hpp:113
stop::StopRule::Ptr stop
Definition TrainingSpec.hpp:116
binarize::Binarization::Ptr binarization
Definition TrainingSpec.hpp:117