Brush C++ API
A flexible interpretable machine learning framework
Loading...
Searching...
No Matches
lexicase.cpp
Go to the documentation of this file.
1#include "lexicase.h"
2
3namespace Brush {
4namespace Sel {
5
6using namespace Brush;
7using namespace Pop;
8using namespace Sel;
9
10template<ProgramType T>
12{
13 this->name = "lexicase";
14 this->survival = surv;
15}
16
17template<ProgramType T>
18vector<size_t> Lexicase<T>::select(Population<T>& pop, int island,
19 const Parameters& params)
20{
21 // this one can be executed in parallel because it is just reading the errors. This
22 // method assumes that the expressions have been fitted previously, and their respective
23 // error vectors are filled
24
25 auto island_pool = pop.get_island_indexes(island);
26
27 // Filter out nullptr individuals (offspring slots)
28 island_pool.erase(
29 std::remove_if(island_pool.begin(), island_pool.end(),
30 [&pop](size_t idx) { return pop.individuals.at(idx) == nullptr; }),
31 island_pool.end()
32 );
33
34 // If all individuals were nullptr, return empty selection
35 if (island_pool.empty())
36 return island_pool;
37
38 // if this is first generation, just return indices to pop
39 if (params.current_gen==0)
40 return island_pool;
41
42 //< number of samples
43 unsigned int N = pop.individuals.at(island_pool.at(0))->error.size();
44
45 //< number of individuals
46 unsigned int P = island_pool.size();
47
48 // define epsilon
49 ArrayXf epsilon = ArrayXf::Zero(N);
50
51 // if output is continuous (per sample!!), use epsilon lexicase.
52 // basically, every classification metric that updates the loss as hit/miss should
53 // not be considered here. If the clf metric updates the reference loss vector
54 // by assigning it float predict probas, then it will need the epsilon lexicase
55 // to work.
56 // The classification scorer names [average_precision_score, roc_auc] are the same
57 // for the binary and multiclassifier
58 if (!params.classification || params.scorer.compare("log")==0
59 || params.scorer.compare("multi_log")==0
60 || params.scorer.compare("average_precision_score")==0
61 || params.scorer.compare("roc_auc")==0 )
62 {
63 // for each sample, calculate epsilon
64 for (int i = 0; i<epsilon.size(); ++i)
65 {
66 VectorXf case_errors(island_pool.size());
67 for (int j = 0; j<island_pool.size(); ++j)
68 {
69 case_errors(j) = pop.individuals.at(island_pool[j])->error(i);
70 }
71
72 // notice that metric used to calculate the error must be a
73 // minimization problem in order for lexicase to work
74 epsilon(i) = mad(case_errors);
75 }
76 }
77 assert(epsilon.size() == N);
78
79 // selection pool
80 vector<size_t> starting_pool;
81 for (int i = 0; i < island_pool.size(); ++i)
82 {
83 starting_pool.push_back(island_pool[i]);
84 }
85 assert(starting_pool.size() == P);
86
87 vector<size_t> selected(P,0); // selected individuals
88
89 for (unsigned int i = 0; i<P; ++i) // selection loop
90 {
91 vector<size_t> cases; // cases (samples)
92 if (params.classification && !params.class_weights.empty())
93 {
94 // NOTE: when calling lexicase, make sure `errors` is from training
95 // data, and not from validation data. This is because the sample
96 // weights indexes are based on train partition
97
98 // for classification problems, weight case selection
99 // by class weights
100 cases.resize(0);
101 vector<size_t> choices(N);
102 std::iota(choices.begin(), choices.end(),0);
103
104 vector<float> sample_weights = params.sample_weights;
105
106 for (unsigned i = 0; i<N; ++i)
107 {
108 vector<size_t> choice_indices(N-i);
109 std::iota(choice_indices.begin(),choice_indices.end(),0);
110
111 size_t idx = *r.select_randomly(
112 choice_indices.begin(), choice_indices.end(),
113 sample_weights.begin(), sample_weights.end());
114
115 cases.push_back(choices.at(idx));
116 choices.erase(choices.begin() + idx);
117
118 sample_weights.erase(sample_weights.begin() + idx);
119 }
120 }
121 else
122 { // otherwise, choose cases randomly
123 cases.resize(N);
124 std::iota(cases.begin(),cases.end(),0);
125 r.shuffle(cases.begin(),cases.end()); // shuffle cases
126 }
127 vector<size_t> pool = starting_pool; // initial pool
128 vector<size_t> winner; // winners
129
130 bool pass = true; // checks pool size and number of cases
131 unsigned int h = 0; // case count
132
133 float epsilon_threshold;
134
135 while(pass){ // main loop
136 epsilon_threshold = 0;
137
138 winner.resize(0); // winners
139 // minimum error on case
140 float minfit = std::numeric_limits<float>::max();
141
142 // get minimum (assuming minization of indiviual errors)
143 for (size_t j = 0; j<pool.size(); ++j)
144 if (pop.individuals.at(pool[j])->error(cases[h]) < minfit)
145 minfit = pop.individuals.at(pool[j])->error(cases[h]);
146
147 // criteria to stay in pool
148 epsilon_threshold = minfit+epsilon[cases[h]];
149
150 // select best
151 for (size_t j = 0; j<pool.size(); ++j)
152 {
153 if (pop.individuals.at(pool[j])->error(cases[h])
154 <= epsilon_threshold)
155 winner.push_back(pool[j]);
156 }
157
158 ++h; // next case
159 // only keep going if needed
160 pass = (winner.size()>1 && h<cases.size());
161
162 if(winner.size() == 0)
163 {
164 if(h >= cases.size())
165 winner.push_back(*r.select_randomly(
166 pool.begin(), pool.end()) );
167 else
168 pass = true;
169 }
170 else
171 pool = winner; // reduce pool to remaining individuals
172 }
173
174 assert(winner.size()>0);
175
176 //if more than one winner, pick randomly
177 selected.at(i) = *r.select_randomly(
178 winner.begin(), winner.end() );
179 }
180
181 if (selected.size() != island_pool.size())
182 {
183 HANDLE_ERROR_THROW("Lexicase did not select correct number of \
184 parents");
185 }
186
187 return selected;
188}
189
190template<ProgramType T>
191vector<size_t> Lexicase<T>::survive(Population<T>& pop, int island,
192 const Parameters& params)
193{
194 /* Lexicase survival */
195 HANDLE_ERROR_THROW("Lexicase survival not implemented");
196 return vector<size_t>();
197}
198
199}
200}
201
vector< size_t > get_island_indexes(int island)
Definition population.h:39
vector< std::shared_ptr< Individual< T > > > individuals
Definition population.h:19
Lexicase selection operator.
Definition lexicase.h:22
Lexicase(bool surv=false)
Definition lexicase.cpp:11
vector< size_t > survive(Population< T > &pop, int island, const Parameters &p)
lexicase survival
Definition lexicase.cpp:191
vector< size_t > select(Population< T > &pop, int island, const Parameters &p)
function returns a set of selected indices from pop
Definition lexicase.cpp:18
#define HANDLE_ERROR_THROW(err)
Definition error.h:27
float mad(const ArrayXf &x)
median absolute deviation
Definition utils.cpp:373
< nsga2 selection operator for getting the front
Definition bandit.cpp:3
vector< float > sample_weights
weights for each sample
Definition params.h:70
bool classification
Definition params.h:75
vector< float > class_weights
weights for each class
Definition params.h:69
string scorer
actual loss function used, determined by error
Definition params.h:66
unsigned int current_gen
Definition params.h:29