MCPcopy Create free account
hub / github.com/EdwardRaff/JSAT / pruneErrorBased

Method pruneErrorBased

JSAT/src/jsat/classifiers/trees/TreePruner.java:166–269  ·  view source on GitHub ↗

@param parent the parent node, or null if there is no parent @param pathFollowed the path from the parent node to the current node @param current the current node to evaluate @param testSet the set of points to estimate error from @param alpha the Confidence @return expected upperbound on errors

(TreeNodeVisitor parent, int pathFollowed, TreeNodeVisitor current, List<DataPointPair<Integer>> testSet, double alpha)

Source from the content-addressed store, hash-verified

164 * @return expected upperbound on errors
165 */
166 private static double pruneErrorBased(TreeNodeVisitor parent, int pathFollowed, TreeNodeVisitor current, List<DataPointPair<Integer>> testSet, double alpha)
167 {
168 //TODO this does a lot of redundant computation. Re-write this code to keep track of where datapoints came from to avoid redudancy.
169 if(current == null || testSet.isEmpty())
170 return 0;
171 else if(current.isLeaf())//return number of errors
172 {
173 int errors = 0;
174 double N = 0;
175 for(DataPointPair<Integer> dpp : testSet)
176 {
177 if(current.localClassify(dpp.getDataPoint()).mostLikely() != dpp.getPair())
178 errors+=dpp.getDataPoint().getWeight();
179 N+=dpp.getDataPoint().getWeight();
180 }
181 return computeBinomialUpperBound(N, alpha, errors);
182 }
183 List<List<DataPointPair<Integer>>> splitSet = new ArrayList<List<DataPointPair<Integer>>>(current.childrenCount());
184 List<DataPointPair<Integer>> hadMissing = new ArrayList<DataPointPair<Integer>>(0);
185 for(int i = 0; i < current.childrenCount(); i++)
186 splitSet.add(new ArrayList<DataPointPair<Integer>>());
187
188 int localErrors = 0;
189 double subTreeScore = 0;
190
191 double N = 0.0;
192 for(DataPointPair<Integer> dpp : testSet)
193 {
194 DataPoint dp = dpp.getDataPoint();
195 if(current.localClassify(dp).mostLikely() != dpp.getPair())
196 localErrors+=dp.getWeight();
197 N += dp.getWeight();
198
199 int path = current.getPath(dp);
200 if(path >= 0)
201 splitSet.get(path).add(dpp);
202 else
203 hadMissing.add(dpp);
204 }
205
206 if(!hadMissing.isEmpty())
207 DecisionStump.distributMissing(splitSet, hadMissing);
208
209 //Find child wich gets the most of the test set as the candidate for sub-tree replacement
210 int maxChildCount = 0;
211 int maxChild = -1;
212 for(int path = 0; path < splitSet.size(); path++)
213 if(!current.isPathDisabled(path))
214 {
215 subTreeScore += pruneErrorBased(current, path, current.getChild(path), splitSet.get(path), alpha);
216
217 if(maxChildCount < splitSet.get(path).size())
218 {
219 maxChildCount = splitSet.get(path).size();
220 maxChild = path;
221 }
222 }
223

Callers 1

pruneMethod · 0.95

Calls 15

getWeightMethod · 0.95
distributMissingMethod · 0.95
classifyMethod · 0.95
mostLikelyMethod · 0.80
sizeMethod · 0.65
isLeafMethod · 0.45
localClassifyMethod · 0.45
getDataPointMethod · 0.45
getPairMethod · 0.45
getWeightMethod · 0.45
childrenCountMethod · 0.45

Tested by

no test coverage detected