gtsam
Loading...
Searching...
No Matches
AlgebraicDecisionTree.h
Go to the documentation of this file.
1/* ----------------------------------------------------------------------------
2
3 * GTSAM Copyright 2010, Georgia Tech Research Corporation,
4 * Atlanta, Georgia 30332-0415
5 * All Rights Reserved
6 * Authors: Frank Dellaert, et al. (see THANKS for the full author list)
7
8 * See LICENSE for the license information
9
10 * -------------------------------------------------------------------------- */
11
18
19#pragma once
20
21#include <gtsam/base/Testable.h>
22#include <gtsam/discrete/DecisionTree-inl.h>
23#include <gtsam/discrete/Ring.h>
24
25#include <iomanip>
26#include <limits>
27#include <map>
28#include <string>
29#include <vector>
30
31namespace gtsam {
32
40 template <typename L>
41 class AlgebraicDecisionTree : public DecisionTree<L, double> {
49 static std::string DefaultFormatter(const L& x) {
50 std::stringstream ss;
51 ss << x;
52 return ss.str();
53 }
54
55 public:
56 using Base = DecisionTree<L, double>;
57
58 AlgebraicDecisionTree(double leaf = 1.0) : Base(leaf) {}
59
61 AlgebraicDecisionTree(const typename Base::NodePtr root) : Base(root) {}
62
63 // Explicitly non-explicit constructor
64 AlgebraicDecisionTree(const Base& add) : Base(add) {}
65
67 AlgebraicDecisionTree(const L& label, double y1, double y2)
68 : Base(label, y1, y2) {}
69
83 AlgebraicDecisionTree(const typename Base::LabelC& labelC, double y1,
84 double y2)
85 : Base(labelC, y1, y2) {}
86
111 AlgebraicDecisionTree //
112 (const std::vector<typename Base::LabelC>& labelCs,
113 const std::vector<double>& ys) {
114 this->root_ =
115 Base::create(labelCs.begin(), labelCs.end(), ys.begin(), ys.end());
116 }
117
126 AlgebraicDecisionTree //
127 (const std::vector<typename Base::LabelC>& labelCs,
128 const std::string& table) {
129 // Convert string to doubles
130 std::vector<double> ys;
131 std::istringstream iss(table);
132 std::copy(std::istream_iterator<double>(iss),
133 std::istream_iterator<double>(), std::back_inserter(ys));
134
135 // now call recursive Create
136 this->root_ =
137 Base::create(labelCs.begin(), labelCs.end(), ys.begin(), ys.end());
138 }
139
147 template <typename Iterator>
148 AlgebraicDecisionTree(Iterator begin, Iterator end, const L& label)
149 : Base(nullptr) {
150 this->root_ = compose(begin, end, label);
151 }
152
159 template <typename M>
160 AlgebraicDecisionTree(const AlgebraicDecisionTree<M>& other,
161 const std::map<M, L>& map) {
162 // Functor for label conversion so we can use `convertFrom`.
163 std::function<L(const M&)> L_of_M = [&map](const M& label) -> L {
164 return map.at(label);
165 };
166 std::function<double(const double&)> op = Ring::id;
167 this->root_ = DecisionTree<L, double>::convertFrom(other.root_, L_of_M, op);
168 }
169
181 template <typename X, typename Func>
183 : Base(other, f) {}
184
186 AlgebraicDecisionTree operator+(const AlgebraicDecisionTree& g) const {
187 return this->apply(g, &Ring::add);
188 }
189
191 AlgebraicDecisionTree operator-() const {
192 return this->apply(&Ring::negate);
193 }
194
196 AlgebraicDecisionTree operator-(const AlgebraicDecisionTree& g) const {
197 return *this + (-g);
198 }
199
201 AlgebraicDecisionTree operator*(const AlgebraicDecisionTree& g) const {
202 return this->apply(g, &Ring::mul);
203 }
204
206 AlgebraicDecisionTree operator/(const AlgebraicDecisionTree& g) const {
207 return this->apply(g, &Ring::div);
208 }
209
211 double sum() const {
212 double sum = 0;
213 auto visitor = [&](double y) { sum += y; };
214 this->visit(visitor);
215 return sum;
216 }
217
224 AlgebraicDecisionTree normalize() const { return (*this) / this->sum(); }
225
227 double min() const {
228 double min = std::numeric_limits<double>::max();
229 auto visitor = [&](double x) { min = x < min ? x : min; };
230 this->visit(visitor);
231 return min;
232 }
233
235 double max() const {
236 // Get the most negative value
237 double max = -std::numeric_limits<double>::max();
238 auto visitor = [&](double x) { max = x > max ? x : max; };
239 this->visit(visitor);
240 return max;
241 }
242
244 AlgebraicDecisionTree sum(const L& label, size_t cardinality) const {
245 return this->combine(label, cardinality, &Ring::add);
246 }
247
249 AlgebraicDecisionTree sum(const typename Base::LabelC& labelC) const {
250 return this->combine(labelC, &Ring::add);
251 }
252
254 void print(const std::string& s = "",
255 const typename Base::LabelFormatter& labelFormatter =
256 &DefaultFormatter) const {
257 auto valueFormatter = [](const double& v) {
258 std::stringstream ss;
259 ss << std::setw(4) << std::setprecision(8) << v;
260 return ss.str();
261 };
262 Base::print(s, labelFormatter, valueFormatter);
263 }
264
266 bool equals(const AlgebraicDecisionTree& other, double tol = 1e-9) const {
267 // lambda for comparison of two doubles upto some tolerance.
268 auto compare = [tol](double a, double b) {
269 return std::abs(a - b) < tol;
270 };
271 return Base::equals(other, compare);
272 }
273 };
274
275template <typename T>
277 : public Testable<AlgebraicDecisionTree<T>> {};
278} // namespace gtsam
Concept check for values that can be used in unit tests.
Real Ring definition.
Global functions in a separate testing namespace.
Definition chartTesting.h:28
A manifold defines a space in which there is a notion of a linear tangent space that can be centered ...
Definition Group.h:37
A helper that implements the traits interface for GTSAM types.
Definition Testable.h:152
An algebraic decision tree fixes the range of a DecisionTree to double.
Definition AlgebraicDecisionTree.h:41
AlgebraicDecisionTree operator/(const AlgebraicDecisionTree &g) const
division
Definition AlgebraicDecisionTree.h:206
AlgebraicDecisionTree sum(const L &label, size_t cardinality) const
sum out variable
Definition AlgebraicDecisionTree.h:244
AlgebraicDecisionTree(const typename Base::LabelC &labelC, double y1, double y2)
Create a new leaf function splitting on a variable.
Definition AlgebraicDecisionTree.h:83
void print(const std::string &s="", const typename Base::LabelFormatter &labelFormatter=&DefaultFormatter) const
print method customized to value type double.
Definition AlgebraicDecisionTree.h:254
AlgebraicDecisionTree sum(const typename Base::LabelC &labelC) const
sum out variable
Definition AlgebraicDecisionTree.h:249
double sum() const
Compute sum of all values.
Definition AlgebraicDecisionTree.h:211
AlgebraicDecisionTree(const AlgebraicDecisionTree< M > &other, const std::map< M, L > &map)
Convert labels from type M to type L.
Definition AlgebraicDecisionTree.h:160
double min() const
Find the minimum values amongst all leaves.
Definition AlgebraicDecisionTree.h:227
AlgebraicDecisionTree normalize() const
Helper method to perform normalization such that all leaves in the tree sum to 1.
Definition AlgebraicDecisionTree.h:224
AlgebraicDecisionTree(const L &label, double y1, double y2)
Create a new leaf function splitting on a variable.
Definition AlgebraicDecisionTree.h:67
AlgebraicDecisionTree(const DecisionTree< L, X > &other, Func f)
Create from an arbitrary DecisionTree<L, X> by operating on it with a functional f.
Definition AlgebraicDecisionTree.h:182
bool equals(const AlgebraicDecisionTree &other, double tol=1e-9) const
Equality method customized to value type double.
Definition AlgebraicDecisionTree.h:266
AlgebraicDecisionTree operator+(const AlgebraicDecisionTree &g) const
sum
Definition AlgebraicDecisionTree.h:186
AlgebraicDecisionTree(const typename Base::NodePtr root)
Constructor which accepts root pointer.
Definition AlgebraicDecisionTree.h:61
double max() const
Find the maximum values amongst all leaves.
Definition AlgebraicDecisionTree.h:235
AlgebraicDecisionTree operator-(const AlgebraicDecisionTree &g) const
subtract
Definition AlgebraicDecisionTree.h:196
AlgebraicDecisionTree operator-() const
negation
Definition AlgebraicDecisionTree.h:191
AlgebraicDecisionTree operator*(const AlgebraicDecisionTree &g) const
product
Definition AlgebraicDecisionTree.h:201
AlgebraicDecisionTree(Iterator begin, Iterator end, const L &label)
Create a range of decision trees, splitting on a single variable.
Definition AlgebraicDecisionTree.h:148
DecisionTree apply(const Unary &op) const
typename Node::Ptr NodePtr
Definition DecisionTree.h:146
NodePtr root_
A DecisionTree just contains the root. TODO(dellaert): make protected.
Definition DecisionTree.h:149
void print(const std::string &s, const LabelFormatter &labelFormatter, const ValueFormatter &valueFormatter) const
DecisionTree combine(const L &label, size_t cardinality, const Binary &op) const
static NodePtr convertFrom(const typename DecisionTree< L, X >::NodePtr &f, std::function< Y(const X &)> Y_of_X)
Convert from a DecisionTree<L, X> to DecisionTree<L, Y>.
Definition DecisionTree-inl.h:745
std::pair< L, size_t > LabelC
Definition DecisionTree.h:80
static NodePtr create(It begin, It end, ValueIt beginY, ValueIt endY)