-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathforest_impl.cpp
More file actions
124 lines (97 loc) · 3.13 KB
/
Copy pathforest_impl.cpp
File metadata and controls
124 lines (97 loc) · 3.13 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#include "forest_impl.h"
#include <ostream>
#include <vector>
namespace NGCForest {
// TTreeNode
TTreeNode::TTreeNode()
: FeatureIndex(0)
, Threshold(0.0)
, Left(TTreeNodePtr())
, Right(TTreeNodePtr())
, Answers()
{
}
size_t TTreeNode::GetFeatureIndex() const {
return FeatureIndex;
}
double TTreeNode::GetThreshold() const {
return Threshold;
}
TTreeNodePtr TTreeNode::GetLeftNode() const {
return Left;
}
TTreeNodePtr TTreeNode::GetRightNode() const {
return Right;
}
const TFeatures &TTreeNode::GetAnswers() const {
return Answers;
}
void TTreeNode::SplitNode(size_t featureIndex, double threshold, TTreeNodePtr left, TTreeNodePtr right) {
FeatureIndex = featureIndex;
Threshold = threshold;
Left = left;
Right = right;
Answers.clear();
}
void TTreeNode::SetAnswers(TFeatures &&answers) {
Answers = std::move(answers);
}
// TTreeImpl
TTreeImpl::TTreeImpl() {
}
TTreeImpl::~TTreeImpl() {
}
const TFeatures &TTreeImpl::Calculate(const TFeatures &features) const {
return DoCalculate(features);
}
void TTreeImpl::Save(std::ostream &fout) const {
DoSave(fout);
}
// TDynamicTreeImpl
TDynamicTreeImpl::TDynamicTreeImpl(TTreeNodePtr root)
: Root(root)
{
}
const TFeatures &TDynamicTreeImpl::DoCalculate(const TFeatures &features) const {
TTreeNodePtr node = Root;
while (!!node->GetLeftNode()) {
size_t idx = node->GetFeatureIndex();
double featureValue = idx < features.size() ? features[idx] : 0.0;
if (featureValue < node->GetThreshold())
node = node->GetLeftNode();
else
node = node->GetRightNode();
}
return node->GetAnswers();
}
void TDynamicTreeImpl::DoSave(std::ostream &fout) const {
// todo: to be done
}
// TObliviousTreeImpl
TObliviousTreeImpl::TObliviousTreeImpl(const std::vector<size_t> &featureIndexes, const std::vector<double> &thresholds, const std::vector<TFeatures> &answers)
: FeatureIndexes(featureIndexes)
, Thresholds(thresholds)
, Answers(answers)
{
}
const TFeatures &TObliviousTreeImpl::DoCalculate(const TFeatures &features) const {
size_t mask = 0;
for (size_t i = 0; i < FeatureIndexes.size(); ++i) {
mask <<= 1;
double val = features[FeatureIndexes[i]];
if (val >= Thresholds[i])
mask |= 1;
}
return Answers[mask];
}
void TObliviousTreeImpl::DoSave(std::ostream &fout) const {
fout << FeatureIndexes.size() << ' ' << Answers.front().size();
for (size_t i = 0; i < FeatureIndexes.size(); ++i)
fout << ' ' << FeatureIndexes[i] << ' ' << Thresholds[i];
for (size_t i = 0; i < Answers.size(); ++i) {
for (size_t j = 0; j < Answers[i].size(); ++j) {
fout << ' ' << Answers[i][j];
}
}
}
} // namespace NGCForest