gtsam
Loading...
Searching...
No Matches
DecisionTree-inl.h
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
19
20#pragma once
21
23
24#include <cassert>
25#include <fstream>
26#include <iterator>
27#include <map>
28#include <optional>
29#include <set>
30#include <sstream>
31#include <stdexcept>
32#include <string>
33#include <vector>
34
35#if defined(__APPLE__)
36#include <TargetConditionals.h> // for TARGET_OS_IPHONE in DecisionTree::dot()
37#endif
38
39namespace gtsam {
40
41 /****************************************************************************/
42 // Node
43 /****************************************************************************/
44#ifdef DT_DEBUG_MEMORY
45 template<typename L, typename Y>
47#endif
48
49 /****************************************************************************/
50 // Leaf
51 /****************************************************************************/
52 template <typename L, typename Y>
53 struct DecisionTree<L, Y>::Leaf : public DecisionTree<L, Y>::Node {
56
58 Leaf() {}
59
62
64 const Y& constant() const {
65 return constant_;
66 }
67
69 bool sameLeaf(const Leaf& q) const override {
70 return constant_ == q.constant_;
71 }
72
74 bool sameLeaf(const Node& q) const override {
75 return (q.isLeaf() && q.sameLeaf(*this));
76 }
77
79 bool equals(const Node& q, const CompareFunc& compare) const override {
80 if (!q.isLeaf()) return false;
81 const Leaf* other = static_cast<const Leaf*>(&q);
82 return compare(this->constant_, other->constant_);
83 }
84
86 void print(const std::string& s, const LabelFormatter& labelFormatter,
87 const ValueFormatter& valueFormatter) const override {
88 std::cout << s << " Leaf " << valueFormatter(constant_) << std::endl;
89 }
90
92 void dot(std::ostream& os, const LabelFormatter& labelFormatter,
93 const ValueFormatter& valueFormatter,
94 bool showZero) const override {
95 const std::string value = valueFormatter(constant_);
96 if (showZero || value.compare("0"))
97 os << "\"" << this->id() << "\" [label=\"" << value
98 << "\", shape=box, rank=sink, height=0.35, fixedsize=true]\n";
99 }
100
102 const Y& operator()(const Assignment<L>& x) const override {
103 return constant_;
104 }
105
107 NodePtr apply(const Unary& op) const override {
108 NodePtr f(new Leaf(op(constant_)));
109 return f;
110 }
111
113 NodePtr apply(const UnaryAssignment& op,
114 const Assignment<L>& assignment) const override {
115 NodePtr f(new Leaf(op(assignment, constant_)));
116 return f;
117 }
118
119 // Apply binary operator "h = f op g" on Leaf node
120 // Note op is not assumed commutative so we need to keep track of order
121 // Simply calls apply on argument to call correct virtual method:
122 // fL.apply_f_op_g(gL) -> gL.apply_g_op_fL(fL) (below)
123 // fL.apply_f_op_g(gC) -> gC.apply_g_op_fL(fL) (Choice)
124 NodePtr apply_f_op_g(const Node& g, const Binary& op) const override {
125 return g.apply_g_op_fL(*this, op);
126 }
127
128 // Applying binary operator to two leaves results in a leaf
129 NodePtr apply_g_op_fL(const Leaf& fL, const Binary& op) const override {
130 // fL op gL
131 NodePtr h(new Leaf(op(fL.constant_, constant_)));
132 return h;
133 }
134
135 // If second argument is a Choice node, call it's apply with leaf as second
136 NodePtr apply_g_op_fC(const Choice& fC, const Binary& op) const override {
137 return fC.apply_fC_op_gL(*this, op); // operand order back to normal
138 }
139
141 NodePtr choose(const L& label, size_t index) const override {
142 return NodePtr(new Leaf(constant()));
143 }
144
145 bool isLeaf() const override { return true; }
146
147 private:
148 using Base = DecisionTree<L, Y>::Node;
149
150#if GTSAM_ENABLE_BOOST_SERIALIZATION
152 friend class boost::serialization::access;
153 template <class ARCHIVE>
154 void serialize(ARCHIVE& ar, const unsigned int /*version*/) {
155 ar & BOOST_SERIALIZATION_BASE_OBJECT_NVP(Base);
156 ar& BOOST_SERIALIZATION_NVP(constant_);
157 }
158#endif
159 }; // Leaf
160
161 /****************************************************************************/
162 // Choice
163 /****************************************************************************/
164 template<typename L, typename Y>
165 struct DecisionTree<L, Y>::Choice: public DecisionTree<L, Y>::Node {
168
170 std::vector<NodePtr> branches_;
171
172 private:
177 size_t allSame_;
178
179 using ChoicePtr = std::shared_ptr<const Choice>;
180
181 public:
184
185 ~Choice() override {
186#ifdef DT_DEBUG_MEMORY
187 std::cout << Node::nrNodes << " destructing (Choice) " << this->id()
188 << std::endl;
189#endif
190 }
191
208 #ifdef GTSAM_DT_MERGING
209
210 static NodePtr Unique(const NodePtr& node) {
211 if (node->isLeaf()) return node; // Leaf node, return as is
212
213 auto choice = std::static_pointer_cast<const Choice>(node);
214 // Choice node, we recurse!
215 // Make non-const copy so we can update
216 auto f = std::make_shared<Choice>(choice->label(), choice->nrChoices());
217
218 // Iterate over all the branches
219 for (const auto& branch : choice->branches_) {
220 f->push_back(Unique(branch));
221 }
222
223 // If all the branches are the same, we can merge them into one
224 if (f->allSame_) {
225 assert(f->branches().size() > 0);
226 auto f0 = std::static_pointer_cast<const Leaf>(f->branches_[0]);
227 return std::make_shared<Leaf>(f0->constant());
228 }
229
230 return f;
231 }
232
233 #else
234
235 static NodePtr Unique(const NodePtr& node) {
236 // No-op when GTSAM_DT_MERGING is not defined
237 return node;
238 }
239
240 #endif
241 bool isLeaf() const override { return false; }
242
244 Choice(const L& label, size_t count) :
245 label_(label), allSame_(true) {
246 branches_.reserve(count);
247 }
248
250 Choice(const Choice& f, const Choice& g, const Binary& op) :
251 allSame_(true) {
252 // Choose what to do based on label
253 if (f.label() > g.label()) {
254 // f higher than g
255 label_ = f.label();
256 size_t count = f.nrChoices();
257 branches_.reserve(count);
258 for (size_t i = 0; i < count; i++) {
259 NodePtr newBranch = f.branches_[i]->apply_f_op_g(g, op);
260 push_back(std::move(newBranch));
261 }
262 } else if (g.label() > f.label()) {
263 // f lower than g
264 label_ = g.label();
265 size_t count = g.nrChoices();
266 branches_.reserve(count);
267 for (size_t i = 0; i < count; i++) {
268 NodePtr newBranch = g.branches_[i]->apply_g_op_fC(f, op);
269 push_back(std::move(newBranch));
270 }
271 } else {
272 // f same level as g
273 label_ = f.label();
274 size_t count = f.nrChoices();
275 branches_.reserve(count);
276 for (size_t i = 0; i < count; i++) {
277 NodePtr newBranch = f.branches_[i]->apply_f_op_g(*g.branches_[i], op);
278 push_back(std::move(newBranch));
279 }
280 }
281 }
282
284 const L& label() const {
285 return label_;
286 }
287
288 size_t nrChoices() const {
289 return branches_.size();
290 }
291
292 const std::vector<NodePtr>& branches() const {
293 return branches_;
294 }
295
296 std::vector<NodePtr>& branches() {
297 return branches_;
298 }
299
301 void push_back(NodePtr&& node) {
302 // allSame_ is restricted to leaf nodes in a decision tree
303 if (allSame_ && !branches_.empty()) {
304 allSame_ = node->sameLeaf(*branches_.back());
305 }
306 branches_.push_back(std::move(node));
307 }
308
310 void print(const std::string& s, const LabelFormatter& labelFormatter,
311 const ValueFormatter& valueFormatter) const override {
312 std::cout << s << " Choice(";
313 std::cout << labelFormatter(label_) << ") " << std::endl;
314 for (size_t i = 0; i < branches_.size(); i++) {
315 branches_[i]->print(s + " " + std::to_string(i), labelFormatter, valueFormatter);
316 }
317 }
318
320 void dot(std::ostream& os, const LabelFormatter& labelFormatter,
321 const ValueFormatter& valueFormatter,
322 bool showZero) const override {
323 const std::string label = labelFormatter(label_);
324 os << "\"" << this->id() << "\" [shape=circle, label=\"" << label
325 << "\"]\n";
326 size_t B = branches_.size();
327 for (size_t i = 0; i < B; i++) {
328 const NodePtr& branch = branches_[i];
329
330 // Check if zero
331 if (!showZero && branch->isLeaf()) {
332 auto leaf = std::static_pointer_cast<const Leaf>(branch);
333 if (valueFormatter(leaf->constant()).compare("0")) continue;
334 }
335
336 os << "\"" << this->id() << "\" -> \"" << branch->id() << "\"";
337 if (B == 2 && i == 0) os << " [style=dashed]";
338 os << std::endl;
339 branch->dot(os, labelFormatter, valueFormatter, showZero);
340 }
341 }
342
344 bool sameLeaf(const Leaf& q) const override {
345 return false;
346 }
347
349 bool sameLeaf(const Node& q) const override {
350 return (q.isLeaf() && q.sameLeaf(*this));
351 }
352
354 bool equals(const Node& q, const CompareFunc& compare) const override {
355 if (q.isLeaf()) return false;
356 const Choice* other = static_cast<const Choice*>(&q);
357 if (this->label_ != other->label_) return false;
358 if (branches_.size() != other->branches_.size()) return false;
359 // we don't care about shared pointers being equal here
360 for (size_t i = 0; i < branches_.size(); i++)
361 if (!(branches_[i]->equals(*(other->branches_[i]), compare)))
362 return false;
363 return true;
364 }
365
367 const Y& operator()(const Assignment<L>& x) const override {
368#ifndef NDEBUG
369 typename Assignment<L>::const_iterator it = x.find(label_);
370 if (it == x.end()) {
371 std::cout << "Trying to find value for " << label_ << std::endl;
372 throw std::invalid_argument(
373 "DecisionTree::operator(): value undefined for a label");
374 }
375#endif
376 size_t index = x.at(label_);
377 NodePtr child = branches_[index];
378 return (*child)(x);
379 }
380
382 Choice(const L& label, const Choice& f, const Unary& op) :
383 label_(label), allSame_(true) {
384 branches_.reserve(f.branches_.size()); // reserve space
385 for (const NodePtr& branch : f.branches_) {
386 push_back(branch->apply(op));
387 }
388 }
389
400 Choice(const L& label, const Choice& f, const UnaryAssignment& op,
401 const Assignment<L>& assignment)
402 : label_(label), allSame_(true) {
403 branches_.reserve(f.branches_.size()); // reserve space
404
405 Assignment<L> assignment_ = assignment;
406
407 for (size_t i = 0; i < f.branches_.size(); i++) {
408 assignment_[label_] = i; // Set assignment for label to i
409
410 const NodePtr branch = f.branches_[i];
411 push_back(branch->apply(op, assignment_));
412
413 // Remove the assignment so we are backtracking
414 auto assignment_it = assignment_.find(label_);
415 assignment_.erase(assignment_it);
416 }
417 }
418
420 NodePtr apply(const Unary& op) const override {
421 auto r = std::make_shared<Choice>(label_, *this, op);
422 return Unique(r);
423 }
424
426 NodePtr apply(const UnaryAssignment& op,
427 const Assignment<L>& assignment) const override {
428 auto r = std::make_shared<Choice>(label_, *this, op, assignment);
429 return Unique(r);
430 }
431
432 // Apply binary operator "h = f op g" on Choice node
433 // Note op is not assumed commutative so we need to keep track of order
434 // Simply calls apply on argument to call correct virtual method:
435 // fC.apply_f_op_g(gL) -> gL.apply_g_op_fC(fC) -> (Leaf)
436 // fC.apply_f_op_g(gC) -> gC.apply_g_op_fC(fC) -> (below)
437 NodePtr apply_f_op_g(const Node& g, const Binary& op) const override {
438 return g.apply_g_op_fC(*this, op);
439 }
440
441 // If second argument of binary op is Leaf node, recurse on branches
442 NodePtr apply_g_op_fL(const Leaf& fL, const Binary& op) const override {
443 auto h = std::make_shared<Choice>(label(), nrChoices());
444 for (auto&& branch : branches_)
445 h->push_back(fL.apply_f_op_g(*branch, op));
446 return Unique(h);
447 }
448
449 // If second argument of binary op is Choice, call constructor
450 NodePtr apply_g_op_fC(const Choice& fC, const Binary& op) const override {
451 auto h = std::make_shared<Choice>(fC, *this, op);
452 return Unique(h);
453 }
454
455 // If second argument of binary op is Leaf
456 template<typename OP>
457 NodePtr apply_fC_op_gL(const Leaf& gL, OP op) const {
458 auto h = std::make_shared<Choice>(label(), nrChoices());
459 for (auto&& branch : branches_)
460 h->push_back(branch->apply_f_op_g(gL, op));
461 return Unique(h);
462 }
463
465 NodePtr choose(const L& label, size_t index) const override {
466 if (label_ == label) return branches_[index]; // choose branch
467
468 // second case, not label of interest, just recurse
469 auto r = std::make_shared<Choice>(label_, branches_.size());
470 for (auto&& branch : branches_) {
471 r->push_back(branch->choose(label, index));
472 }
473
474 return Unique(r);
475 }
476
477 private:
478 using Base = DecisionTree<L, Y>::Node;
479
480#if GTSAM_ENABLE_BOOST_SERIALIZATION
482 friend class boost::serialization::access;
483 template <class ARCHIVE>
484 void serialize(ARCHIVE& ar, const unsigned int /*version*/) {
485 ar & BOOST_SERIALIZATION_BASE_OBJECT_NVP(Base);
486 ar& BOOST_SERIALIZATION_NVP(label_);
487 ar& BOOST_SERIALIZATION_NVP(branches_);
488 ar& BOOST_SERIALIZATION_NVP(allSame_);
489 }
490#endif
491 }; // Choice
492
493 /****************************************************************************/
494 // DecisionTree
495 /****************************************************************************/
496 template <typename L, typename Y>
498
499 template<typename L, typename Y>
500 DecisionTree<L, Y>::DecisionTree(const NodePtr& root) :
501 root_(root) {}
502
503 /****************************************************************************/
504 template<typename L, typename Y>
506 root_ = NodePtr(new Leaf(y));
507 }
508
509 /****************************************************************************/
510 template <typename L, typename Y>
511 DecisionTree<L, Y>::DecisionTree(const L& label, const Y& y1, const Y& y2) {
512 auto a = std::make_shared<Choice>(label, 2);
513 NodePtr l1(new Leaf(y1)), l2(new Leaf(y2));
514 a->push_back(std::move(l1));
515 a->push_back(std::move(l2));
516 root_ = Choice::Unique(std::move(a));
517 }
518
519 /****************************************************************************/
520 template <typename L, typename Y>
521 DecisionTree<L, Y>::DecisionTree(const LabelC& labelC, const Y& y1,
522 const Y& y2) {
523 if (labelC.second != 2) throw std::invalid_argument(
524 "DecisionTree: binary constructor called with non-binary label");
525 auto a = std::make_shared<Choice>(labelC.first, 2);
526 NodePtr l1(new Leaf(y1)), l2(new Leaf(y2));
527 a->push_back(std::move(l1));
528 a->push_back(std::move(l2));
529 root_ = Choice::Unique(std::move(a));
530 }
531 /****************************************************************************/
532 template<typename L, typename Y>
533 DecisionTree<L, Y>::DecisionTree(const std::vector<LabelC>& labelCs,
534 const std::vector<Y>& ys) {
535 // call recursive Create
536 root_ = create(labelCs.begin(), labelCs.end(), ys.begin(), ys.end());
537 }
538
539 /****************************************************************************/
540 template<typename L, typename Y>
541 DecisionTree<L, Y>::DecisionTree(const std::vector<LabelC>& labelCs,
542 const std::string& table) {
543 // Convert std::string to values of type Y
544 std::vector<Y> ys;
545 std::istringstream iss(table);
546 copy(std::istream_iterator<Y>(iss), std::istream_iterator<Y>(),
547 back_inserter(ys));
548
549 // now call recursive Create
550 root_ = create(labelCs.begin(), labelCs.end(), ys.begin(), ys.end());
551 }
552
553 /****************************************************************************/
554 template<typename L, typename Y>
555 template<typename Iterator> DecisionTree<L, Y>::DecisionTree(
556 Iterator begin, Iterator end, const L& label) {
557 root_ = compose(begin, end, label);
558 }
559
560 /****************************************************************************/
561 template<typename L, typename Y>
563 const DecisionTree& f0, const DecisionTree& f1) {
564 const std::vector<DecisionTree> functions{f0, f1};
565 root_ = compose(functions.begin(), functions.end(), label);
566 }
567
568 /****************************************************************************/
569 template <typename L, typename Y>
571 DecisionTree&& other) noexcept
572 : root_(std::move(other.root_)) {
573 // Apply the unary operation directly to each leaf in the tree
574 if (root_) {
575 // Define a helper function to traverse and apply the operation
576 struct ApplyUnary {
577 const Unary& op;
578 void operator()(typename DecisionTree<L, Y>::NodePtr& node) const {
579 if (node->isLeaf()) {
580 // Apply the unary operation to the leaf's constant value
581 auto leaf = std::static_pointer_cast<Leaf>(node);
582 leaf->constant_ = op(leaf->constant_);
583 } else {
584 // Recurse into the choice branches
585 auto choice = std::static_pointer_cast<Choice>(node);
586 for (NodePtr& branch : choice->branches()) {
587 (*this)(branch);
588 }
589 }
590 }
591 };
592
593 ApplyUnary applyUnary{op};
594 applyUnary(root_);
595 }
596 // Reset the other tree's root to nullptr to avoid dangling references
597 other.root_ = nullptr;
598 }
599
600 /****************************************************************************/
601 template <typename L, typename Y>
602 template <typename X, typename Func>
604 Func Y_of_X) {
605 root_ = convertFrom<X>(other.root_, Y_of_X);
606 }
607
608 /****************************************************************************/
609 template <typename L, typename Y>
610 template <typename M, typename X, typename Func>
612 const std::map<M, L>& map, Func Y_of_X) {
613 auto L_of_M = [&map](const M& label) -> L { return map.at(label); };
614 root_ = convertFrom<M, X>(other.root_, L_of_M, Y_of_X);
615 }
616
617 /****************************************************************************/
618 // Called by two constructors above.
619 // Takes a label and a corresponding range of decision trees, and creates a
620 // new decision tree. However, the order of the labels needs to be respected,
621 // so we cannot just create a root Choice node on the label: if the label is
622 // not the highest label, we need a complicated/ expensive recursive call.
623 template <typename L, typename Y>
624 template <typename Iterator>
625 typename DecisionTree<L, Y>::NodePtr DecisionTree<L, Y>::compose(
626 Iterator begin, Iterator end, const L& label) {
627 // find highest label among branches
628 std::optional<L> highestLabel;
629 size_t nrChoices = 0;
630 for (Iterator it = begin; it != end; it++) {
631 if (it->root_->isLeaf())
632 continue;
633 auto c = std::static_pointer_cast<const Choice>(it->root_);
634 if (!highestLabel || c->label() > *highestLabel) {
635 highestLabel = c->label();
636 nrChoices = c->nrChoices();
637 }
638 }
639
640 // if label is already in correct order, just put together a choice on label
641 if (!nrChoices || !highestLabel || label > *highestLabel) {
642 auto choiceOnLabel = std::make_shared<Choice>(label, end - begin);
643 for (Iterator it = begin; it != end; it++) {
644 NodePtr root = it->root_;
645 choiceOnLabel->push_back(std::move(root));
646 }
647 // If no reordering, no need to call Choice::Unique
648 return choiceOnLabel;
649 } else {
650 // Set up a new choice on the highest label
651 auto choiceOnHighestLabel =
652 std::make_shared<Choice>(*highestLabel, nrChoices);
653 // now, for all possible values of highestLabel
654 for (size_t index = 0; index < nrChoices; index++) {
655 // make a new set of functions for composing by iterating over the given
656 // functions, and selecting the appropriate branch.
657 std::vector<DecisionTree> functions;
658 for (Iterator it = begin; it != end; it++) {
659 // by restricting the input functions to value i for labelBelow
660 DecisionTree chosen = it->choose(*highestLabel, index);
661 functions.push_back(chosen);
662 }
663 // We then recurse, for all values of the highest label
664 NodePtr fi = compose(functions.begin(), functions.end(), label);
665 choiceOnHighestLabel->push_back(std::move(fi));
666 }
667 return choiceOnHighestLabel;
668 }
669 }
670
671 /****************************************************************************/
672 // "build" is a bit of a complicated thing, but very useful.
673 // It takes a range of labels and a corresponding range of values,
674 // and builds a decision tree, as follows:
675 // - if there is only one label, creates a choice node with values in leaves
676 // - otherwise, it evenly splits up the range of values and creates a tree for
677 // each sub-range, and assigns that tree to first label's choices
678 // Example:
679 // build([B A],[1 2 3 4]) would call
680 // build([A],[1 2])
681 // build([A],[3 4])
682 // and produce
683 // B=0
684 // A=0: 1
685 // A=1: 2
686 // B=1
687 // A=0: 3
688 // A=1: 4
689 // Note, through the magic of "compose", create([A B],[1 3 2 4]) will produce
690 // exactly the same tree as above: the highest label is always the root.
691 // However, it will be *way* faster if labels are given highest to lowest.
692 template<typename L, typename Y>
693 template<typename It, typename ValueIt>
695 It begin, It end, ValueIt beginY, ValueIt endY) {
696 // get crucial counts
697 size_t nrChoices = begin->second;
698 size_t size = endY - beginY;
699
700 // Find the next key to work on
701 It labelC = begin + 1;
702 if (labelC == end) {
703 // Base case: only one key left
704 // Create a simple choice node with values as leaves.
705 if (size != nrChoices) {
706 std::cout << "Trying to create DD on " << begin->first << std::endl;
707 std::cout << "DecisionTree::create: expected " << nrChoices
708 << " values but got " << size << " instead" << std::endl;
709 throw std::invalid_argument("DecisionTree::create invalid argument");
710 }
711 auto choice = std::make_shared<Choice>(begin->first, endY - beginY);
712 for (ValueIt y = beginY; y != endY; y++) {
713 choice->push_back(NodePtr(new Leaf(*y)));
714 }
715 return choice;
716 }
717
718 // Recursive case: perform "Shannon expansion"
719 // Creates one tree (i.e.,function) for each choice of current key
720 // by calling create recursively, and then puts them all together.
721 std::vector<DecisionTree> functions;
722 functions.reserve(nrChoices);
723 size_t split = size / nrChoices;
724 for (size_t i = 0; i < nrChoices; i++, beginY += split) {
725 NodePtr f = build<It, ValueIt>(labelC, end, beginY, beginY + split);
726 functions.emplace_back(f);
727 }
728 return compose(functions.begin(), functions.end(), begin->first);
729 }
730
731 /****************************************************************************/
732 // Top-level factory method, which takes a range of labels and a corresponding
733 // range of values, and creates a decision tree.
734 template<typename L, typename Y>
735 template<typename It, typename ValueIt>
737 It begin, It end, ValueIt beginY, ValueIt endY) {
738 auto node = build(begin, end, beginY, endY);
739 return Choice::Unique(node);
740 }
741
742 /****************************************************************************/
743 template <typename L, typename Y>
744 template <typename X>
746 const typename DecisionTree<L, X>::NodePtr& f,
747 std::function<Y(const X&)> Y_of_X) {
748 using LXLeaf = typename DecisionTree<L, X>::Leaf;
749 using LXChoice = typename DecisionTree<L, X>::Choice;
750
751 // If leaf, apply unary conversion "op" and create a unique leaf.
752 if (f->isLeaf()) {
753 auto leaf = std::static_pointer_cast<LXLeaf>(f);
754 return NodePtr(new Leaf(Y_of_X(leaf->constant())));
755 }
756
757 // Now a Choice!
758 auto choice = std::static_pointer_cast<const LXChoice>(f);
759
760 // Create a new Choice node with the same label
761 auto newChoice = std::make_shared<Choice>(choice->label(), choice->nrChoices());
762
763 // Convert each branch recursively
764 for (auto&& branch : choice->branches()) {
765 newChoice->push_back(convertFrom<X>(branch, Y_of_X));
766 }
767
768 return Choice::Unique(newChoice);
769 }
770
771 /****************************************************************************/
772 template <typename L, typename Y>
773 template <typename M, typename X>
775 const typename DecisionTree<M, X>::NodePtr& f,
776 std::function<L(const M&)> L_of_M, std::function<Y(const X&)> Y_of_X) {
777 using LY = DecisionTree<L, Y>;
778 using MXLeaf = typename DecisionTree<M, X>::Leaf;
779 using MXChoice = typename DecisionTree<M, X>::Choice;
780
781 // If leaf, apply unary conversion "op" and create a unique leaf.
782 if (f->isLeaf()) {
783 auto leaf = std::static_pointer_cast<const MXLeaf>(f);
784 return NodePtr(new Leaf(Y_of_X(leaf->constant())));
785 }
786
787 // Now is Choice!
788 auto choice = std::static_pointer_cast<const MXChoice>(f);
789
790 // get new label
791 const M oldLabel = choice->label();
792 const L newLabel = L_of_M(oldLabel);
793
794 // Shannon expansion in this context involves:
795 // 1. Creating separate subtrees (functions) for each possible value of the new label.
796 // 2. Combining these subtrees using the 'compose' method, which implements the expansion.
797 // This approach guarantees that the resulting tree maintains the correct variable ordering
798 // based on the new labels (L) after translation from the old labels (M).
799 // Simply creating a Choice node here would not work because it wouldn't account for the
800 // potentially new ordering of variables resulting from the label translation,
801 // which is crucial for maintaining consistency and efficiency in the converted tree.
802 std::vector<LY> functions;
803 for (auto&& branch : choice->branches()) {
804 functions.emplace_back(convertFrom<M, X>(branch, L_of_M, Y_of_X));
805 }
806 return Choice::Unique(
807 LY::compose(functions.begin(), functions.end(), newLabel));
808 }
809
810 /****************************************************************************/
821 template <typename L, typename Y>
822 struct Visit {
823 using F = std::function<void(const Y&)>;
824 explicit Visit(F f) : f(f) {}
825 F f;
826
828 void operator()(const typename DecisionTree<L, Y>::NodePtr& node) const {
829 using Leaf = typename DecisionTree<L, Y>::Leaf;
830 using Choice = typename DecisionTree<L, Y>::Choice;
831
832 if (node->isLeaf()) {
833 auto leaf = std::static_pointer_cast<const Leaf>(node);
834 return f(leaf->constant());
835 }
836
837 auto choice = std::static_pointer_cast<const Choice>(node);
838 for (auto&& branch : choice->branches()) (*this)(branch); // recurse!
839 }
840 };
841
842 template <typename L, typename Y>
843 template <typename Func>
844 void DecisionTree<L, Y>::visit(Func f) const {
846 visit(root_);
847 }
848
849 /****************************************************************************/
859 template <typename L, typename Y>
860 struct VisitLeaf {
861 using F = std::function<void(const typename DecisionTree<L, Y>::Leaf&)>;
862 explicit VisitLeaf(F f) : f(f) {}
863 F f;
864
866 void operator()(const typename DecisionTree<L, Y>::NodePtr& node) const {
867 using Leaf = typename DecisionTree<L, Y>::Leaf;
868 using Choice = typename DecisionTree<L, Y>::Choice;
869
870 if (node->isLeaf()) {
871 auto leaf = std::static_pointer_cast<const Leaf>(node);
872 return f(*leaf);
873 }
874
875 auto choice = std::static_pointer_cast<const Choice>(node);
876 for (auto&& branch : choice->branches()) (*this)(branch); // recurse!
877 }
878 };
879
880 template <typename L, typename Y>
881 template <typename Func>
882 void DecisionTree<L, Y>::visitLeaf(Func f) const {
884 visit(root_);
885 }
886
887 /****************************************************************************/
894 template <typename L, typename Y>
895 struct VisitWith {
896 using F = std::function<void(const Assignment<L>&, const Y&)>;
897 explicit VisitWith(F f) : f(f) {}
899 F f;
900
902 void operator()(const typename DecisionTree<L, Y>::NodePtr& node) {
903 using Leaf = typename DecisionTree<L, Y>::Leaf;
904 using Choice = typename DecisionTree<L, Y>::Choice;
905
906 if (node->isLeaf()) {
907 auto leaf = std::static_pointer_cast<const Leaf>(node);
908 return f(assignment, leaf->constant());
909 }
910
911
912
913 auto choice = std::static_pointer_cast<const Choice>(node);
914 for (size_t i = 0; i < choice->nrChoices(); i++) {
915 assignment[choice->label()] = i; // Set assignment for label to i
916
917 (*this)(choice->branches()[i]); // recurse!
918
919 // Remove the choice so we are backtracking
920 auto choice_it = assignment.find(choice->label());
921 assignment.erase(choice_it);
922 }
923 }
924 };
925
926 template <typename L, typename Y>
927 template <typename Func>
928 void DecisionTree<L, Y>::visitWith(Func f) const {
930 visit(root_);
931 }
932
933 /****************************************************************************/
934 template <typename L, typename Y>
936 size_t total = 0;
937 visit([&total](const Y& node) { total += 1; });
938 return total;
939 }
940
941 /****************************************************************************/
942 // fold is just done with a visit
943 template <typename L, typename Y>
944 template <typename Func, typename X>
945 X DecisionTree<L, Y>::fold(Func f, X x0) const {
946 visit([&](const Y& y) { x0 = f(y, x0); });
947 return x0;
948 }
949
950 /****************************************************************************/
964 template <typename L, typename Y>
965 std::set<L> DecisionTree<L, Y>::labels() const {
966 std::set<L> unique;
967 auto f = [&](const Assignment<L>& assignment, const Y&) {
968 for (auto&& kv : assignment) {
969 unique.insert(kv.first);
970 }
971 };
972 visitWith(f);
973 return unique;
974 }
975
976/****************************************************************************/
977 template <typename L, typename Y>
978 bool DecisionTree<L, Y>::equals(const DecisionTree& other,
979 const CompareFunc& compare) const {
980 return root_->equals(*other.root_, compare);
981 }
982
983 template <typename L, typename Y>
984 void DecisionTree<L, Y>::print(const std::string& s,
985 const LabelFormatter& labelFormatter,
986 const ValueFormatter& valueFormatter) const {
987 root_->print(s, labelFormatter, valueFormatter);
988 }
989
990 template<typename L, typename Y>
992 return root_->equals(*other.root_);
993 }
994
995 /****************************************************************************/
996 template<typename L, typename Y>
998 if (root_ == nullptr)
999 throw std::invalid_argument(
1000 "DecisionTree::operator() called on empty tree");
1001 return root_->operator ()(x);
1002 }
1003
1004 /****************************************************************************/
1005 template<typename L, typename Y>
1007 // It is unclear what should happen if tree is empty:
1008 if (empty()) {
1009 throw std::runtime_error(
1010 "DecisionTree::apply(unary op) undefined for empty tree.");
1011 }
1012 return DecisionTree(root_->apply(op));
1013 }
1014
1015 /****************************************************************************/
1017 template <typename L, typename Y>
1019 const UnaryAssignment& op) const {
1020 // It is unclear what should happen if tree is empty:
1021 if (empty()) {
1022 throw std::runtime_error(
1023 "DecisionTree::apply(unary op) undefined for empty tree.");
1024 }
1025 Assignment<L> assignment;
1026 return DecisionTree(root_->apply(op, assignment));
1027 }
1028
1029 /****************************************************************************/
1030 template<typename L, typename Y>
1032 const Binary& op) const {
1033 // It is unclear what should happen if either tree is empty:
1034 if (empty() || g.empty()) {
1035 throw std::runtime_error(
1036 "DecisionTree::apply(binary op) undefined for empty trees.");
1037 }
1038 // apply the operaton on the root of both diagrams
1039 NodePtr h = root_->apply_f_op_g(*g.root_, op);
1040 // create a new class with the resulting root "h"
1041 DecisionTree result(h);
1042 return result;
1043 }
1044
1045 /****************************************************************************/
1046 // The way this works:
1047 // We have an ADT, picture it as a tree.
1048 // At a certain depth, we have a branch on "label".
1049 // The function "choose(label,index)" will return a tree of one less depth,
1050 // where there is no more branch on "label": only the subtree under that
1051 // branch point corresponding to the value "index" is left instead.
1052 // The function below get all these smaller trees and "ops" them together.
1053 // This implements marginalization in Darwiche09book, pg 330
1054 template<typename L, typename Y>
1056 size_t cardinality, const Binary& op) const {
1057 DecisionTree result = choose(label, 0);
1058 for (size_t index = 1; index < cardinality; index++) {
1059 DecisionTree chosen = choose(label, index);
1060 result = result.apply(chosen, op);
1061 }
1062 return result;
1063 }
1064
1065 /****************************************************************************/
1066 template <typename L, typename Y>
1067 void DecisionTree<L, Y>::dot(std::ostream& os,
1068 const LabelFormatter& labelFormatter,
1069 const ValueFormatter& valueFormatter,
1070 bool showZero) const {
1071 os << "digraph G {\n";
1072 root_->dot(os, labelFormatter, valueFormatter, showZero);
1073 os << " [ordering=out]}" << std::endl;
1074 }
1075
1076 template <typename L, typename Y>
1077 void DecisionTree<L, Y>::dot(const std::string& name,
1078 const LabelFormatter& labelFormatter,
1079 const ValueFormatter& valueFormatter,
1080 bool showZero) const {
1081 std::ofstream os((name + ".dot").c_str());
1082 dot(os, labelFormatter, valueFormatter, showZero);
1083#if defined(__APPLE__) && TARGET_OS_IPHONE
1084 // iOS marks std::system() unavailable, breaking the build. The PDF
1085 // rendering is a debug convenience over the always-emitted .dot file;
1086 // callers who want a PDF on iOS-targeted builds can run
1087 // `dot -Tpdf <name>.dot -o <name>.pdf` themselves on a host with a shell.
1088#else
1089 int result =
1090 system(("dot -Tpdf " + name + ".dot -o " + name + ".pdf >& /dev/null")
1091 .c_str());
1092 if (result == -1)
1093 throw std::runtime_error("DecisionTree::dot system call failed");
1094#endif
1095 }
1096
1097 template <typename L, typename Y>
1098 std::string DecisionTree<L, Y>::dot(const LabelFormatter& labelFormatter,
1099 const ValueFormatter& valueFormatter,
1100 bool showZero) const {
1101 std::stringstream ss;
1102 dot(ss, labelFormatter, valueFormatter, showZero);
1103 return ss.str();
1104 }
1105
1106 /******************************************************************************/
1107 template <typename L, typename Y>
1108 template <typename A, typename B>
1109 std::pair<DecisionTree<L, A>, DecisionTree<L, B>> DecisionTree<L, Y>::split(
1110 std::function<std::pair<A, B>(const Y&)> AB_of_Y) const {
1111 using AB = std::pair<A, B>;
1112 const DecisionTree<L, AB> ab(*this, AB_of_Y);
1113 const DecisionTree<L, A> a(ab, [](const AB& p) { return p.first; });
1114 const DecisionTree<L, B> b(ab, [](const AB& p) { return p.second; });
1115 return {a, b};
1116 }
1117
1118 /******************************************************************************/
1119
1120 } // namespace gtsam
Decision Tree for use in DiscreteFactors.
Global functions in a separate testing namespace.
Definition chartTesting.h:28
double dot(const V1 &a, const V2 &b)
Dot product.
Definition Vector.h:191
An assignment from labels to value index (size_t).
Definition Assignment.h:37
Definition DecisionTree-inl.h:53
NodePtr choose(const L &label, size_t index) const override
choose a branch, create new memory !
Definition DecisionTree-inl.h:141
const Y & operator()(const Assignment< L > &x) const override
evaluate
Definition DecisionTree-inl.h:102
NodePtr apply(const UnaryAssignment &op, const Assignment< L > &assignment) const override
Apply unary operator with assignment.
Definition DecisionTree-inl.h:113
bool equals(const Node &q, const CompareFunc &compare) const override
equality up to tolerance
Definition DecisionTree-inl.h:79
Y constant_
constant stored in this leaf
Definition DecisionTree-inl.h:55
void print(const std::string &s, const LabelFormatter &labelFormatter, const ValueFormatter &valueFormatter) const override
print
Definition DecisionTree-inl.h:86
NodePtr apply(const Unary &op) const override
apply unary operator
Definition DecisionTree-inl.h:107
Leaf(const Y &constant)
Constructor from constant.
Definition DecisionTree-inl.h:61
bool sameLeaf(const Leaf &q) const override
Leaf-Leaf equality.
Definition DecisionTree-inl.h:69
void dot(std::ostream &os, const LabelFormatter &labelFormatter, const ValueFormatter &valueFormatter, bool showZero) const override
Write graphviz format to stream os.
Definition DecisionTree-inl.h:92
Leaf()
Default constructor for serialization.
Definition DecisionTree-inl.h:58
bool sameLeaf(const Node &q) const override
polymorphic equality: is q a leaf and is it the same as this leaf?
Definition DecisionTree-inl.h:74
const Y & constant() const
Return the constant.
Definition DecisionTree-inl.h:64
Definition DecisionTree-inl.h:165
NodePtr apply(const Unary &op) const override
apply unary operator.
Definition DecisionTree-inl.h:420
void push_back(NodePtr &&node)
add a branch: TODO merge into constructor
Definition DecisionTree-inl.h:301
Choice(const L &label, const Choice &f, const UnaryAssignment &op, const Assignment< L > &assignment)
Constructor which accepts a UnaryAssignment op and the corresponding assignment.
Definition DecisionTree-inl.h:400
const L & label() const
Return the label of this choice node.
Definition DecisionTree-inl.h:284
void print(const std::string &s, const LabelFormatter &labelFormatter, const ValueFormatter &valueFormatter) const override
print (as a tree).
Definition DecisionTree-inl.h:310
static NodePtr Unique(const NodePtr &node)
Merge branches with equal leaf values for every choice node in a decision tree.
Definition DecisionTree-inl.h:235
NodePtr apply(const UnaryAssignment &op, const Assignment< L > &assignment) const override
Apply unary operator with assignment.
Definition DecisionTree-inl.h:426
L label_
the label of the variable on which we split
Definition DecisionTree-inl.h:167
bool sameLeaf(const Node &q) const override
polymorphic equality: if q is a leaf, could be...
Definition DecisionTree-inl.h:349
Choice(const Choice &f, const Choice &g, const Binary &op)
Construct from applying binary op to two Choice nodes.
Definition DecisionTree-inl.h:250
std::vector< NodePtr > branches_
The children of this Choice node.
Definition DecisionTree-inl.h:170
Choice()
Default constructor for serialization.
Definition DecisionTree-inl.h:183
const Y & operator()(const Assignment< L > &x) const override
evaluate
Definition DecisionTree-inl.h:367
Choice(const L &label, size_t count)
Constructor, given choice label and mandatory expected branch count.
Definition DecisionTree-inl.h:244
NodePtr choose(const L &label, size_t index) const override
choose a branch, recursively
Definition DecisionTree-inl.h:465
Choice(const L &label, const Choice &f, const Unary &op)
Construct from applying unary op to a Choice node.
Definition DecisionTree-inl.h:382
void dot(std::ostream &os, const LabelFormatter &labelFormatter, const ValueFormatter &valueFormatter, bool showZero) const override
output to graphviz (as a a graph)
Definition DecisionTree-inl.h:320
bool sameLeaf(const Leaf &q) const override
Choice-Leaf equality: always false.
Definition DecisionTree-inl.h:344
bool equals(const Node &q, const CompareFunc &compare) const override
equality
Definition DecisionTree-inl.h:354
Functor performing depth-first visit to each leaf with the leaf value as the argument.
Definition DecisionTree-inl.h:822
F f
folding function object.
Definition DecisionTree-inl.h:825
void operator()(const typename DecisionTree< L, Y >::NodePtr &node) const
Do a depth-first visit on the tree rooted at node.
Definition DecisionTree-inl.h:828
Visit(F f)
Construct from folding function.
Definition DecisionTree-inl.h:824
Functor performing depth-first visit to each leaf with the Leaf object passed as an argument.
Definition DecisionTree-inl.h:860
VisitLeaf(F f)
Construct from folding function.
Definition DecisionTree-inl.h:862
void operator()(const typename DecisionTree< L, Y >::NodePtr &node) const
Do a depth-first visit on the tree rooted at node.
Definition DecisionTree-inl.h:866
F f
folding function object.
Definition DecisionTree-inl.h:863
Functor performing depth-first visit to each leaf with the leaf's Assignment<L> and value passed as a...
Definition DecisionTree-inl.h:895
VisitWith(F f)
Construct from folding function.
Definition DecisionTree-inl.h:897
Assignment< L > assignment
Assignment, mutating through recursion.
Definition DecisionTree-inl.h:898
void operator()(const typename DecisionTree< L, Y >::NodePtr &node)
Do a depth-first visit on the tree rooted at node.
Definition DecisionTree-inl.h:902
F f
folding function object.
Definition DecisionTree-inl.h:899
a decision tree is a function from assignments to values.
Definition DecisionTree.h:62
DecisionTree apply(const Unary &op) const
apply Unary operation "op" to f
Definition DecisionTree-inl.h:1006
DecisionTree choose(const L &label, size_t index) const
create a new function where value(label)==index It's like "restrict" in Darwiche09book pg329,...
Definition DecisionTree.h:391
typename Node::Ptr NodePtr
---------------------— Node base class ------------------------—
Definition DecisionTree.h:146
static NodePtr build(It begin, It end, ValueIt beginY, ValueIt endY)
Internal recursive function to create from keys, cardinalities, and Y values.
Definition DecisionTree-inl.h:694
std::set< L > labels() const
Retrieve all unique labels as a set.
Definition DecisionTree-inl.h:965
bool empty() const
Check if tree is empty.
Definition DecisionTree.h:290
void visit(Func f) const
Visit all leaves in depth-first fashion.
Definition DecisionTree-inl.h:844
void visitLeaf(Func f) const
Visit all leaves in depth-first fashion.
Definition DecisionTree-inl.h:882
std::function< Y(const Y &)> Unary
Handy typedefs for unary and binary function types.
Definition DecisionTree.h:75
X fold(Func f, X x0) const
Fold a binary function over the tree, returning accumulator.
Definition DecisionTree-inl.h:945
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
GTSAM-style print.
Definition DecisionTree-inl.h:984
DecisionTree combine(const L &label, size_t cardinality, const Binary &op) const
combine subtrees on key with binary operation "op"
Definition DecisionTree-inl.h:1055
void visitWith(Func f) const
Visit all leaves in depth-first fashion.
Definition DecisionTree-inl.h:928
const Y & operator()(const Assignment< L > &x) const
evaluate
Definition DecisionTree-inl.h:997
void dot(std::ostream &os, const LabelFormatter &labelFormatter, const ValueFormatter &valueFormatter, bool showZero=true) const
output to graphviz format, stream version
Definition DecisionTree-inl.h:1067
std::pair< DecisionTree< L, A >, DecisionTree< L, B > > split(std::function< std::pair< A, B >(const Y &)> AB_of_Y) const
Convert into two trees with value types A and B.
Definition DecisionTree-inl.h:1109
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
bool operator==(const DecisionTree &q) const
equality
Definition DecisionTree-inl.h:991
std::pair< L, size_t > LabelC
A label annotated with cardinality.
Definition DecisionTree.h:80
size_t nrLeaves() const
Return the number of leaves in the tree.
Definition DecisionTree-inl.h:935
static NodePtr create(It begin, It end, ValueIt beginY, ValueIt endY)
Internal helper function to create a tree from keys, cardinalities, and Y values.
Definition DecisionTree-inl.h:736
DecisionTree()
Default constructor (for serialization).
Definition DecisionTree-inl.h:497
---------------------— Node base class ------------------------—
Definition DecisionTree.h:87