gtsam
Loading...
Searching...
No Matches
DecisionTreeFactor.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
24#include <gtsam/discrete/Ring.h>
25
26#include <memory>
27#include <string>
28#include <utility>
29#include <vector>
30
31namespace gtsam {
32
34 class HybridValues;
35
41 class GTSAM_EXPORT DecisionTreeFactor : public DiscreteFactor,
42 public AlgebraicDecisionTree<Key> {
43 public:
44 // typedefs needed to play nice with gtsam
45 typedef DecisionTreeFactor This;
47 typedef std::shared_ptr<DecisionTreeFactor> shared_ptr;
48 typedef AlgebraicDecisionTree<Key> ADT;
49
50 // Needed since we have definitions in both DiscreteFactor and DecisionTree
51 using Base::Binary;
52 using Base::Unary;
53 using Base::UnaryAssignment;
54
57
60
62 DecisionTreeFactor(const DiscreteKeys& keys, const ADT& potentials);
63
84 const std::vector<double>& table);
85
104 DecisionTreeFactor(const DiscreteKeys& keys, const std::string& table);
105
107 template <class SOURCE>
108 DecisionTreeFactor(const DiscreteKey& key, SOURCE table)
109 : DecisionTreeFactor(DiscreteKeys{key}, table) {}
110
112 DecisionTreeFactor(const DiscreteKey& key, const std::vector<double>& row)
113 : DecisionTreeFactor(DiscreteKeys{key}, row) {}
114
116 explicit DecisionTreeFactor(const DiscreteConditional& c);
117
121
123 bool equals(const DiscreteFactor& other, double tol = 1e-9) const override;
124
125 // print
126 void print(
127 const std::string& s = "DecisionTreeFactor:\n",
128 const KeyFormatter& formatter = DefaultKeyFormatter) const override;
129
133
136 virtual double evaluate(const Assignment<Key>& values) const override {
137 return ADT::operator()(values);
138 }
139
141 using DiscreteFactor::operator();
142
144 double error(const DiscreteValues& values) const override;
145
160 virtual DiscreteFactor::shared_ptr multiply(
161 const DiscreteFactor::shared_ptr& f) const override;
162
164 DiscreteFactor::shared_ptr operator*(double s) const override {
165 return std::make_shared<DecisionTreeFactor>(
166 apply([s](const double& a) { return Ring::mul(a, s); }));
167 }
168
171 return apply(f, Ring::mul);
172 }
173
174 static double safe_div(const double& a, const double& b);
175
186 return apply(f, safe_div);
187 }
188
191 const DiscreteFactor::shared_ptr& f) const override;
192
194 DecisionTreeFactor toDecisionTreeFactor() const override { return *this; }
195
197 TableFactor toTableFactor() const override;
198
200 using ADT::sum;
201
203 DiscreteFactor::shared_ptr sum(size_t nrFrontals) const override {
204 return combine(nrFrontals, Ring::add);
205 }
206
209 return combine(keys, Ring::add);
210 }
211
213 double max() const override { return ADT::max(); };
214
216 DiscreteFactor::shared_ptr max(size_t nrFrontals) const override {
217 return combine(nrFrontals, Ring::max);
218 }
219
222 return combine(keys, Ring::max);
223 }
224
227 const DiscreteValues& assignment) const override;
228
232
237 DecisionTreeFactor apply(Unary op) const;
238
244 DecisionTreeFactor apply(UnaryAssignment op) const;
245
251 DecisionTreeFactor apply(const DecisionTreeFactor& f, Binary op) const;
252
259 shared_ptr combine(size_t nrFrontals, Binary op) const;
260
267 shared_ptr combine(const Ordering& keys, Binary op) const;
268
270 std::vector<std::pair<DiscreteValues, double>> enumerate() const;
271
273 std::vector<double> probabilities() const;
274
284 double computeThreshold(const size_t N) const;
285
304 DecisionTreeFactor prune(size_t maxNrAssignments) const;
305
310 uint64_t nrValues() const override { return nrLeaves(); }
311
315
317 void dot(std::ostream& os,
318 const KeyFormatter& keyFormatter = DefaultKeyFormatter,
319 bool showZero = true) const;
320
322 void dot(const std::string& name,
323 const KeyFormatter& keyFormatter = DefaultKeyFormatter,
324 bool showZero = true) const;
325
327 std::string dot(const KeyFormatter& keyFormatter = DefaultKeyFormatter,
328 bool showZero = true) const;
329
337 std::string markdown(const KeyFormatter& keyFormatter = DefaultKeyFormatter,
338 const Names& names = {}) const override;
339
347 std::string html(const KeyFormatter& keyFormatter = DefaultKeyFormatter,
348 const Names& names = {}) const override;
349
353
358 double error(const HybridValues& values) const override;
359
361
362 private:
363#if GTSAM_ENABLE_BOOST_SERIALIZATION
365 friend class boost::serialization::access;
366 template <class ARCHIVE>
367 void serialize(ARCHIVE& ar, const unsigned int /*version*/) {
368 ar& BOOST_SERIALIZATION_BASE_OBJECT_NVP(Base);
369 ar& BOOST_SERIALIZATION_BASE_OBJECT_NVP(ADT);
370 }
371#endif
372 };
373
374// traits
375template <>
376struct traits<DecisionTreeFactor> : public Testable<DecisionTreeFactor> {};
377} // namespace gtsam
specialized key for discrete variables
Algebraic Decision Trees.
Real Ring definition.
std::pair< Key, size_t > DiscreteKey
Key type for discrete variables.
Definition DiscreteKey.h:38
Global functions in a separate testing namespace.
Definition chartTesting.h:28
KeyFormatter DefaultKeyFormatter
Assign default key formatter.
Definition Key.cpp:30
void print(const Matrix &A, const string &s, ostream &stream)
print without optional string, must specify cout yourself
Definition Matrix.cpp:143
string markdown(const DiscreteValues &values, const KeyFormatter &keyFormatter, const DiscreteValues::Names &names)
Free version of markdown.
Definition DiscreteValues.cpp:155
std::function< std::string(Key)> KeyFormatter
Typedef for a function to format a key, i.e. to convert it to a string.
Definition Key.h:35
DecisionTree< L, Y > apply(const DecisionTree< L, Y > &f, const typename DecisionTree< L, Y >::Unary &op)
free versions of apply
Definition DecisionTree.h:467
double dot(const V1 &a, const V2 &b)
Dot product.
Definition Vector.h:191
A manifold defines a space in which there is a notion of a linear tangent space that can be centered ...
Definition Group.h:37
Template to create a binary predicate.
Definition Testable.h:112
A helper that implements the traits interface for GTSAM types.
Definition Testable.h:152
double max() const
Definition AlgebraicDecisionTree.h:235
An assignment from labels to value index (size_t).
Definition Assignment.h:37
const Y & operator()(const Assignment< L > &x) const
evaluate
Definition DecisionTree-inl.h:997
size_t nrLeaves() const
Return the number of leaves in the tree.
Definition DecisionTree-inl.h:935
A discrete probabilistic factor.
Definition DecisionTreeFactor.h:42
DecisionTreeFactor operator*(const DecisionTreeFactor &f) const override
multiply two factors
Definition DecisionTreeFactor.h:170
virtual double evaluate(const Assignment< Key > &values) const override
Calculate probability for given values, is just look up in AlgebraicDecisionTree.
Definition DecisionTreeFactor.h:136
DecisionTreeFactor apply(Unary op) const
Apply unary operator (*this) "op" f.
Definition DecisionTreeFactor.cpp:135
shared_ptr combine(size_t nrFrontals, Binary op) const
Combine frontal variables using binary operator "op".
Definition DecisionTreeFactor.cpp:170
DiscreteFactor::shared_ptr max(size_t nrFrontals) const override
Create new factor by maximizing over all values with the same separator.
Definition DecisionTreeFactor.h:216
DiscreteFactor::shared_ptr max(const Ordering &keys) const override
Create new factor by maximizing over all values with the same separator.
Definition DecisionTreeFactor.h:221
double max() const override
Find the maximum value in the factor.
Definition DecisionTreeFactor.h:213
DiscreteFactor Base
Typedef to base class.
Definition DecisionTreeFactor.h:46
uint64_t nrValues() const override
Get the number of non-zero values contained in this factor.
Definition DecisionTreeFactor.h:310
DiscreteFactor::shared_ptr sum(const Ordering &keys) const override
Create new factor by summing all values with the same separator values.
Definition DecisionTreeFactor.h:208
DecisionTreeFactor(const DiscreteKey &key, SOURCE table)
Single-key specialization.
Definition DecisionTreeFactor.h:108
DiscreteFactor::shared_ptr sum(size_t nrFrontals) const override
Create new factor by summing all values with the same separator values.
Definition DecisionTreeFactor.h:203
DecisionTreeFactor toDecisionTreeFactor() const override
Convert into a decision tree.
Definition DecisionTreeFactor.h:194
DiscreteFactor::shared_ptr operator*(double s) const override
multiply with a scalar
Definition DecisionTreeFactor.h:164
DecisionTreeFactor()
Default constructor for I/O.
Definition DecisionTreeFactor.cpp:34
DecisionTreeFactor operator/(const DecisionTreeFactor &f) const
Divide by factor f (safely).
Definition DecisionTreeFactor.h:185
DecisionTreeFactor(const DiscreteKey &key, const std::vector< double > &row)
Single-key specialization, with vector of doubles.
Definition DecisionTreeFactor.h:112
Discrete Conditional Density Derives from DecisionTreeFactor.
Definition DiscreteConditional.h:40
Base class for discrete probabilistic factors The most general one is the derived DecisionTreeFactor.
Definition DiscreteFactor.h:41
std::shared_ptr< DiscreteFactor > shared_ptr
shared_ptr to this class
Definition DiscreteFactor.h:46
DiscreteFactor()
Default constructor creates empty factor.
Definition DiscreteFactor.h:65
DiscreteKeys is a set of keys that can be assembled using the & operator.
Definition DiscreteKey.h:41
A map from keys to values.
Definition DiscreteValues.h:34
A discrete probabilistic factor optimized for sparsity.
Definition TableFactor.h:51
HybridValues represents a collection of DiscreteValues and VectorValues.
Definition HybridValues.h:37
const KeyVector & keys() const
Access the factor's involved variable keys.
Definition Factor.h:143
Definition Ordering.h:33