gtsam
Loading...
Searching...
No Matches
BayesTreeMarginalizationHelper.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
19
20// \callgraph
21#pragma once
22
23#include <unordered_map>
24#include <unordered_set>
25#include <deque>
28#include <gtsam/base/debug.h>
29#include "gtsam/dllexport.h"
30#include <stdexcept>
31#include <string>
32
33namespace gtsam {
34
38template <typename BayesTree>
40
41public:
42 using Clique = typename BayesTree::Clique;
43 using sharedClique = typename BayesTree::sharedClique;
44
68 static std::unordered_set<Key>
70 const BayesTree& bayesTree,
71 const KeyVector& marginalizableKeys) {
72 const bool debug = ISDEBUG("BayesTreeMarginalizationHelper");
73
74 std::unordered_set<const Clique*> additionalCliques =
75 gatherAdditionalCliquesToReEliminate(bayesTree, marginalizableKeys);
76
77 std::unordered_set<Key> additionalKeys;
78 for (const Clique* clique : additionalCliques) {
79 addCliqueToKeySet(clique, &additionalKeys);
80 }
81
82 if (debug) {
83 std::cout << "BayesTreeMarginalizationHelper: Additional keys to re-eliminate: ";
84 for (const Key& key : additionalKeys) {
85 std::cout << DefaultKeyFormatter(key) << " ";
86 }
87 std::cout << std::endl;
88 }
89
90 return additionalKeys;
91 }
92
93 protected:
99 static std::unordered_set<const Clique*>
101 const BayesTree& bayesTree,
102 const KeyVector& marginalizableKeys) {
103 std::unordered_set<const Clique*> additionalCliques;
104 std::unordered_set<Key> marginalizableKeySet(
105 marginalizableKeys.begin(), marginalizableKeys.end());
106 CachedSearch cachedSearch;
107
108 // Check each clique that contains a marginalizable key
109 for (const Clique* clique :
110 getCliquesContainingKeys(bayesTree, marginalizableKeySet)) {
111 if (additionalCliques.count(clique)) {
112 // The clique has already been visited. This can happen when an
113 // ancestor of the current clique also contain some marginalizable
114 // varaibles and it's processed beore the current.
115 continue;
116 }
117
118 if (needsReelimination(clique, marginalizableKeySet, &cachedSearch)) {
119 // Add the current clique
120 additionalCliques.insert(clique);
121
122 // Then add the dependent cliques
123 gatherDependentCliques(clique, marginalizableKeySet, &additionalCliques,
124 &cachedSearch);
125 }
126 }
127 return additionalCliques;
128 }
129
137 static std::unordered_set<const Clique*> getCliquesContainingKeys(
138 const BayesTree& bayesTree,
139 const std::unordered_set<Key>& keysOfInterest) {
140 std::unordered_set<const Clique*> cliques;
141 const auto& nodes = bayesTree.nodes();
142 for (const Key& key : keysOfInterest) {
143 // One lookup, reused for both the check and the insert. exists() is a
144 // count() and bayesTree[key] is an at(), so testing then indexing would
145 // search the node map twice per key; bayesTree[key] also returns the
146 // shared_ptr by value, which costs a reference count round trip.
147 const auto node = nodes.find(key);
148 if (node == nodes.end()) {
149 // Otherwise the lookup fails inside bayesTree[key], naming neither the
150 // key nor this helper, with text that varies by build: "map::at"
151 // without TBB, "Index out of requested size range" with it. A variable
152 // that no factor references is never eliminated, so it has no clique.
153 throw std::out_of_range(
154 "BayesTreeMarginalizationHelper: key '" + DefaultKeyFormatter(key) +
155 "' has no clique in the Bayes tree.");
156 }
157 cliques.insert(node->second.get());
158 }
159 return cliques;
160 }
161
166 std::unordered_map<const Clique*, bool> wholeMarginalizableCliques;
167 std::unordered_map<const Clique*, bool> wholeMarginalizableSubtrees;
168 };
169
176 const Clique* clique,
177 const std::unordered_set<Key>& marginalizableKeys,
178 CachedSearch* cache) {
179 auto it = cache->wholeMarginalizableCliques.find(clique);
180 if (it != cache->wholeMarginalizableCliques.end()) {
181 return it->second;
182 } else {
183 bool ret = true;
184 for (Key key : clique->conditional()->frontals()) {
185 if (!marginalizableKeys.count(key)) {
186 ret = false;
187 break;
188 }
189 }
190 cache->wholeMarginalizableCliques.insert({clique, ret});
191 return ret;
192 }
193 }
194
201 const Clique* subtree,
202 const std::unordered_set<Key>& marginalizableKeys,
203 CachedSearch* cache) {
204 auto it = cache->wholeMarginalizableSubtrees.find(subtree);
205 if (it != cache->wholeMarginalizableSubtrees.end()) {
206 return it->second;
207 } else {
208 bool ret = true;
209 if (isWholeCliqueMarginalizable(subtree, marginalizableKeys, cache)) {
210 for (const sharedClique& child : subtree->children) {
211 if (!isWholeSubtreeMarginalizable(child.get(), marginalizableKeys, cache)) {
212 ret = false;
213 break;
214 }
215 }
216 } else {
217 ret = false;
218 }
219 cache->wholeMarginalizableSubtrees.insert({subtree, ret});
220 return ret;
221 }
222 }
223
233 const Clique* clique,
234 const std::unordered_set<Key>& marginalizableKeys,
235 CachedSearch* cache) {
236 bool hasNonMarginalizableAhead = false;
237
238 // Check each frontal variable in order
239 for (Key key : clique->conditional()->frontals()) {
240 if (marginalizableKeys.count(key)) {
241 // If we've seen non-marginalizable variables before this one,
242 // we need to reeliminate
243 if (hasNonMarginalizableAhead) {
244 return true;
245 }
246
247 // Check if any child depends on this marginalizable key and the
248 // subtree rooted at that child contains non-marginalizables.
249 for (const sharedClique& child : clique->children) {
250 if (hasDependency(child.get(), key) &&
251 !isWholeSubtreeMarginalizable(child.get(), marginalizableKeys, cache)) {
252 return true;
253 }
254 }
255 } else {
256 hasNonMarginalizableAhead = true;
257 }
258 }
259 return false;
260 }
261
270 const Clique* rootClique,
271 const std::unordered_set<Key>& marginalizableKeys,
272 std::unordered_set<const Clique*>* additionalCliques,
273 CachedSearch* cache) {
274 std::vector<const Clique*> dependentChildren;
275 dependentChildren.reserve(rootClique->children.size());
276 for (const sharedClique& child : rootClique->children) {
277 if (additionalCliques->count(child.get())) {
278 // This child has already been visited. This can happen if the
279 // child itself contains a marginalizable variable and it's
280 // processed before the current rootClique.
281 continue;
282 }
283 if (hasDependency(child.get(), marginalizableKeys)) {
284 dependentChildren.push_back(child.get());
285 }
286 }
288 dependentChildren, marginalizableKeys, additionalCliques, cache);
289 }
290
295 const std::vector<const Clique*>& dependentChildren,
296 const std::unordered_set<Key>& marginalizableKeys,
297 std::unordered_set<const Clique*>* additionalCliques,
298 CachedSearch* cache) {
299 std::deque<const Clique*> descendants(
300 dependentChildren.begin(), dependentChildren.end());
301 while (!descendants.empty()) {
302 const Clique* descendant = descendants.front();
303 descendants.pop_front();
304
305 // If the subtree rooted at this descendant contains non-marginalizables,
306 // it must lie on a path from the root clique to a clique containing
307 // non-marginalizables at the leaf side.
308 if (!isWholeSubtreeMarginalizable(descendant, marginalizableKeys, cache)) {
309 additionalCliques->insert(descendant);
310
311 // Add children of the current descendant to the set descendants.
312 for (const sharedClique& child : descendant->children) {
313 if (additionalCliques->count(child.get())) {
314 // This child has already been visited.
315 continue;
316 } else {
317 descendants.push_back(child.get());
318 }
319 }
320 }
321 }
322 }
323
330 static void addCliqueToKeySet(
331 const Clique* clique,
332 std::unordered_set<Key>* additionalKeys) {
333 for (Key key : clique->conditional()->frontals()) {
334 additionalKeys->insert(key);
335 }
336 }
337
345 static bool hasDependency(
346 const Clique* clique, Key key) {
347 auto& conditional = clique->conditional();
348 if (std::find(conditional->beginParents(),
349 conditional->endParents(), key)
350 != conditional->endParents()) {
351 return true;
352 } else {
353 return false;
354 }
355 }
356
360 static bool hasDependency(
361 const Clique* clique, const std::unordered_set<Key>& keys) {
362 auto& conditional = clique->conditional();
363 for (auto it = conditional->beginParents();
364 it != conditional->endParents(); ++it) {
365 if (keys.count(*it)) {
366 return true;
367 }
368 }
369
370 return false;
371 }
372};
373// BayesTreeMarginalizationHelper
374
375}
Global debugging flags.
Base class for cliques of a BayesTree.
Bayes Tree is a tree of cliques of a Bayes Chain.
Global functions in a separate testing namespace.
Definition chartTesting.h:28
KeyFormatter DefaultKeyFormatter
Assign default key formatter.
Definition Key.cpp:30
FastVector< Key > KeyVector
Define collection type once and for all - also used in wrappers.
Definition Key.h:91
std::uint64_t Key
Integer nonlinear key type.
Definition types.h:43
Bayes tree.
Definition BayesTree.h:77
std::shared_ptr< Clique > sharedClique
Shared pointer to a clique.
Definition BayesTree.h:84
const Nodes & nodes() const
Return nodes.
Definition BayesTree.h:157
CLIQUE Clique
The clique type, normally BayesTreeClique.
Definition BayesTree.h:83
This class provides helper functions for marginalizing variables from a Bayes Tree.
Definition BayesTreeMarginalizationHelper.h:39
static bool hasDependency(const Clique *clique, Key key)
Check if the clique depends on the given key.
Definition BayesTreeMarginalizationHelper.h:345
static std::unordered_set< const Clique * > gatherAdditionalCliquesToReEliminate(const BayesTree &bayesTree, const KeyVector &marginalizableKeys)
This function identifies cliques that need to be re-eliminated before performing marginalization.
Definition BayesTreeMarginalizationHelper.h:100
static bool hasDependency(const Clique *clique, const std::unordered_set< Key > &keys)
Check if the clique depends on any of the given keys.
Definition BayesTreeMarginalizationHelper.h:360
static std::unordered_set< Key > gatherAdditionalKeysToReEliminate(const BayesTree &bayesTree, const KeyVector &marginalizableKeys)
This function identifies variables that need to be re-eliminated before performing marginalization.
Definition BayesTreeMarginalizationHelper.h:69
static bool isWholeSubtreeMarginalizable(const Clique *subtree, const std::unordered_set< Key > &marginalizableKeys, CachedSearch *cache)
Check if all variables in the subtree are marginalizable.
Definition BayesTreeMarginalizationHelper.h:200
static void gatherDependentCliques(const Clique *rootClique, const std::unordered_set< Key > &marginalizableKeys, std::unordered_set< const Clique * > *additionalCliques, CachedSearch *cache)
Gather all dependent nodes that lie on a path from the root clique to a clique containing a non-margi...
Definition BayesTreeMarginalizationHelper.h:269
static void gatherDependentCliquesFromChildren(const std::vector< const Clique * > &dependentChildren, const std::unordered_set< Key > &marginalizableKeys, std::unordered_set< const Clique * > *additionalCliques, CachedSearch *cache)
A helper function for the above gatherDependentCliques().
Definition BayesTreeMarginalizationHelper.h:294
static void addCliqueToKeySet(const Clique *clique, std::unordered_set< Key > *additionalKeys)
Add all frontal variables from a clique to a key set.
Definition BayesTreeMarginalizationHelper.h:330
static std::unordered_set< const Clique * > getCliquesContainingKeys(const BayesTree &bayesTree, const std::unordered_set< Key > &keysOfInterest)
Gather the cliques containing any of the given keys.
Definition BayesTreeMarginalizationHelper.h:137
static bool needsReelimination(const Clique *clique, const std::unordered_set< Key > &marginalizableKeys, CachedSearch *cache)
Check if a clique contains variables that need reelimination due to elimination ordering conflicts.
Definition BayesTreeMarginalizationHelper.h:232
static bool isWholeCliqueMarginalizable(const Clique *clique, const std::unordered_set< Key > &marginalizableKeys, CachedSearch *cache)
Check if all variables in the clique are marginalizable.
Definition BayesTreeMarginalizationHelper.h:175
A struct to cache the results of the below two functions.
Definition BayesTreeMarginalizationHelper.h:165