aGrUM 3.2.0
a C++ library for (probabilistic) graphical models
barrenNodesFinder.cpp
Go to the documentation of this file.
1/****************************************************************************
2 * This file is part of the aGrUM/pyAgrum library. *
3 * *
4 * Copyright (c) 2005-2026 by *
5 * - Pierre-Henri WUILLEMIN(_at_LIP6) *
6 * - Christophe GONZALES(_at_AMU) *
7 * *
8 * The aGrUM/pyAgrum library is free software; you can redistribute it *
9 * and/or modify it under the terms of either : *
10 * *
11 * - the GNU Lesser General Public License as published by *
12 * the Free Software Foundation, either version 3 of the License, *
13 * or (at your option) any later version, *
14 * - the MIT license (MIT), *
15 * - or both in dual license, as here. *
16 * *
17 * (see https://agrum.gitlab.io/articles/dual-licenses-lgplv3mit.html) *
18 * *
19 * This aGrUM/pyAgrum library is distributed in the hope that it will be *
20 * useful, but WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, *
21 * INCLUDING BUT NOT LIMITED TO THE WARRANTIES MERCHANTABILITY or FITNESS *
22 * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE *
23 * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER *
24 * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, *
25 * ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR *
26 * OTHER DEALINGS IN THE SOFTWARE. *
27 * *
28 * See LICENCES for more details. *
29 * *
30 * SPDX-FileCopyrightText: Copyright 2005-2026 *
31 * - Pierre-Henri WUILLEMIN(_at_LIP6) *
32 * - Christophe GONZALES(_at_AMU) *
33 * SPDX-License-Identifier: LGPL-3.0-or-later OR MIT *
34 * *
35 * Contact : info_at_agrum_dot_org *
36 * homepage : http://agrum.gitlab.io *
37 * gitlab : https://gitlab.com/agrumery/agrum *
38 * *
39 ****************************************************************************/
40
41
46#include <limits>
47
49
50#ifdef GUM_NO_INLINE
52#endif // GUM_NO_INLINE
53
54namespace gum {
55
56 // Out-of-line on purpose (not INLINE, see barrenNodesFinder_inl.h): under
57 // MSVC, dllexport on a class forces eager, non-weak emission of every
58 // inline-defined special member in *every* TU that includes the header --
59 // colliding (LNK2005) with the leaf .pyd's own local reinstantiation of
60 // LazyPropagation<double> (GUM_NO_EXTERN_TEMPLATE_CLASS), which calls
61 // straight into these.
62
64 BarrenNodesFinder::BarrenNodesFinder(const DAG* dag) : _dag_(dag) { // for debugging purposes
65 GUM_CONSTRUCTOR(BarrenNodesFinder);
66 }
67
71 _target_nodes_(from._target_nodes_) { // for debugging purposes
72 GUM_CONS_CPY(BarrenNodesFinder);
73 }
74
77 _dag_(from._dag_), _observed_nodes_(from._observed_nodes_),
78 _target_nodes_(from._target_nodes_) {
79 // for debugging purposes
80 GUM_CONS_MOV(BarrenNodesFinder);
81 }
82
84 BarrenNodesFinder::~BarrenNodesFinder() { // for debugging purposes
85 GUM_DESTRUCTOR(BarrenNodesFinder);
86 }
87
90 if (this != &from) {
91 _dag_ = from._dag_;
94 }
95 return *this;
96 }
97
100 if (this != &from) {
101 _dag_ = from._dag_;
102 _observed_nodes_ = from._observed_nodes_;
103 _target_nodes_ = from._target_nodes_;
104 }
105 return *this;
106 }
107
110 // assign a mark to all the nodes
111 // and mark all the observed nodes and their ancestors as non-barren
113 {
114 for (const auto node: *_dag_)
115 mark.insert(node, 0); // for the moment, 0 = possibly barren
116
117 // mark all the observed nodes and their ancestors as non barren
118 // std::numeric_limits<unsigned int>::max () will be necessarily non
119 // barren
120 // later on
121 Sequence< NodeId > observed_anc(_dag_->size());
122 const Size non_barren = std::numeric_limits< Size >::max();
123 for (const auto node: *_observed_nodes_)
124 observed_anc.insert(node);
125 for (Idx i = 0; i < observed_anc.size(); ++i) {
126 const NodeId node = observed_anc[i];
127 if (!mark[node]) {
128 mark[node] = non_barren;
129 for (const auto par: _dag_->parents(node)) {
130 if (!mark[par] && !observed_anc.exists(par)) { observed_anc.insert(par); }
131 }
132 }
133 }
134 }
135
136 // create the data structure that will contain the result of the
137 // method. By default, we assume that, for each pair of adjacent cliques,
138 // all
139 // the nodes that do not belong to their separator are possibly barren and,
140 // by sweeping the dag, we will remove the nodes that were determined
141 // above as non-barren. Structure result will assign to each (ordered) pair
142 // of adjacent cliques its set of barren nodes.
144 for (const auto& edge: junction_tree.edges()) {
145 const NodeSet& separator = junction_tree.separator(edge);
146
147 NodeSet non_barren1 = junction_tree.clique(edge.first());
148 for (auto iter = non_barren1.beginSafe(); iter != non_barren1.endSafe(); ++iter) {
149 if (mark[*iter] || separator.exists(*iter)) { non_barren1.erase(iter); }
150 }
151 result.insert(Arc(edge.first(), edge.second()), std::move(non_barren1));
152
153 NodeSet non_barren2 = junction_tree.clique(edge.second());
154 for (auto iter = non_barren2.beginSafe(); iter != non_barren2.endSafe(); ++iter) {
155 if (mark[*iter] || separator.exists(*iter)) { non_barren2.erase(iter); }
156 }
157 result.insert(Arc(edge.second(), edge.first()), std::move(non_barren2));
158 }
159
160 // for each node in the DAG, indicate which are the arcs in the result
161 // structure whose separator contain it: the separators are actually the
162 // targets of the queries.
163 NodeProperty< ArcSet > node2arc;
164 for (const auto node: *_dag_)
165 node2arc.insert(node, ArcSet());
166 for (const auto& elt: result) {
167 const Arc& arc = elt.first;
168 if (!result[arc].empty()) { // no need to further process cliques
169 const NodeSet& separator = // with no barren nodes
170 junction_tree.separator(Edge(arc.tail(), arc.head()));
171
172 for (const auto node: separator) {
173 node2arc[node].insert(arc);
174 }
175 }
176 }
177
178 // To determine the set of non-barren nodes w.r.t. a given single node
179 // query, we rely on the fact that those are precisely all the ancestors of
180 // this single node. To mutualize the computations, we will thus sweep the
181 // DAG from top to bottom and exploit the fact that the set of ancestors of
182 // the child of a given node A contain the ancestors of A. Therefore, we
183 // will
184 // determine sets of paths in the DAG and, for each path, compute the set of
185 // its barren nodes from the source to the destination of the path. The
186 // optimal set of paths, i.e., that which will minimize computations, is
187 // obtained by solving a "minimum path cover in directed acyclic graphs".
188 // But
189 // such an algorithm is too costly for the gain we can get from it, so we
190 // will
191 // rely on a simple heuristics.
192
193 // To compute the heuristics, we proceed as follows:
194 // 1/ we mark to 1 all the nodes that are ancestors of at least one (key)
195 // node
196 // with a non-empty arcset in node2arc and we extract from those the
197 // roots, i.e., those nodes whose set of parents, if any, have all been
198 // identified as non-barren by being marked as
199 // std::numeric_limits<unsigned int>::max (). Such nodes are
200 // thus the top of the graph to sweep.
201 // 2/ create a copy of the subgraph of the DAG w.r.t. the 1-marked nodes
202 // and, for each node, if the node has several parents and children,
203 // keep only one arc from one of the parents to the child with the
204 // smallest
205 // number of parents, and try to create a matching between parents and
206 // children and add one arc for each edge of this matching. This will
207 // allow
208 // us to create distinct paths in the DAG. Whenever a child has no more
209 // parents, it becomes the root of a new path.
210 // 3/ the sweeping will be performed from the roots of all these paths.
211
212 // perform step 1/
213 NodeSet path_roots;
214 {
215 List< NodeId > nodes_to_mark;
216 for (const auto& elt: node2arc) {
217 if (!elt.second.empty()) { // only process nodes with assigned arcs
218 nodes_to_mark.insert(elt.first);
219 }
220 }
221 while (!nodes_to_mark.empty()) {
222 NodeId node = nodes_to_mark.front();
223 nodes_to_mark.popFront();
224
225 if (!mark[node]) { // mark the node and all its ancestors
226 mark[node] = 1;
227 Size nb_par = 0;
228 for (auto par: _dag_->parents(node)) {
229 Size parent_mark = mark[par];
230 if (parent_mark != std::numeric_limits< Size >::max()) {
231 ++nb_par;
232 if (parent_mark == 0) { nodes_to_mark.insert(par); }
233 }
234 }
235
236 if (nb_par == 0) { path_roots.insert(node); }
237 }
238 }
239 }
240
241 // perform step 2/
242 DAG sweep_dag = *_dag_;
243 for (const auto node: *_dag_) { // keep only nodes marked with 1
244 if (mark[node] != 1) { sweep_dag.eraseNode(node); }
245 }
246 for (const auto node: sweep_dag) {
247 const Size nb_parents = sweep_dag.parents(node).size();
248 const Size nb_children = sweep_dag.children(node).size();
249 if ((nb_parents > 1) || (nb_children > 1)) {
250 // perform the matching
251 const auto& parents = sweep_dag.parents(node);
252
253 // if there is no child, remove all the parents except the first one
254 if (nb_children == 0) {
255 auto iter_par = parents.beginSafe();
256 for (++iter_par; iter_par != parents.endSafe(); ++iter_par) {
257 sweep_dag.eraseArc(Arc(*iter_par, node));
258 }
259 } else {
260 // find the child with the smallest number of parents
261 const auto& children = sweep_dag.children(node);
262 NodeId smallest_child = 0;
263 Size smallest_nb_par = std::numeric_limits< Size >::max();
264 for (const auto child: children) {
265 const auto new_nb = sweep_dag.parents(child).size();
266 if (new_nb < smallest_nb_par) {
267 smallest_nb_par = new_nb;
268 smallest_child = child;
269 }
270 }
271
272 // if there is no parent, just keep the link with the smallest child
273 // and remove all the other arcs
274 if (nb_parents == 0) {
275 for (auto iter = children.beginSafe(); iter != children.endSafe(); ++iter) {
276 if (*iter != smallest_child) {
277 if (sweep_dag.parents(*iter).size() == 1) { path_roots.insert(*iter); }
278 sweep_dag.eraseArc(Arc(node, *iter));
279 }
280 }
281 } else {
282 auto nb_match = Size(std::min(nb_parents, nb_children) - 1);
283 auto iter_par = parents.beginSafe();
284 ++iter_par; // skip the first parent, whose arc with node will
285 // remain
286 auto iter_child = children.beginSafe();
287 for (Idx i = 0; i < nb_match; ++i, ++iter_par, ++iter_child) {
288 if (*iter_child == smallest_child) { ++iter_child; }
289 sweep_dag.addArc(*iter_par, *iter_child);
290 sweep_dag.eraseArc(Arc(*iter_par, node));
291 sweep_dag.eraseArc(Arc(node, *iter_child));
292 }
293 for (; iter_par != parents.endSafe(); ++iter_par) {
294 sweep_dag.eraseArc(Arc(*iter_par, node));
295 }
296 for (; iter_child != children.endSafe(); ++iter_child) {
297 if (*iter_child != smallest_child) {
298 if (sweep_dag.parents(*iter_child).size() == 1) { path_roots.insert(*iter_child); }
299 sweep_dag.eraseArc(Arc(node, *iter_child));
300 }
301 }
302 }
303 }
304 }
305 }
306
307 // step 3: sweep the paths from the roots of sweep_dag
308 // here, the idea is that, for each path of sweep_dag, the mark we put
309 // to the ancestors is a given number, say N, that increases from path
310 // to path. Hence, for a given path, all the nodes that are marked with a
311 // number at least as high as N are non-barren, the others being barren.
312 Idx mark_id = 2;
313 for (NodeId path: path_roots) {
314 // perform the sweeping from the path
315 while (true) {
316 // mark all the ancestors of the node
317 List< NodeId > to_mark{path};
318 while (!to_mark.empty()) {
319 NodeId node = to_mark.front();
320 to_mark.popFront();
321 if (mark[node] < mark_id) {
322 mark[node] = mark_id;
323 for (const auto par: _dag_->parents(node)) {
324 if (mark[par] < mark_id) { to_mark.insert(par); }
325 }
326 }
327 }
328
329 // now, get all the arcs that contained node "path" in their separator.
330 // this node acts as a query target and, therefore, its ancestors
331 // shall be non-barren.
332 const ArcSet& arcs = node2arc[path];
333 for (const auto& arc: arcs) {
334 NodeSet& barren = result[arc];
335 for (auto iter = barren.beginSafe(); iter != barren.endSafe(); ++iter) {
336 if (mark[*iter] >= mark_id) {
337 // this indicates a non-barren node
338 barren.erase(iter);
339 }
340 }
341 }
342
343 // go to the next sweeping node
344 const NodeSet& sweep_children = sweep_dag.children(path);
345 if (sweep_children.size()) {
346 path = *(sweep_children.begin());
347 } else {
348 // here, the path has ended, so we shall go to the next path
349 ++mark_id;
350 break;
351 }
352 }
353 }
354
355 return result;
356 }
357
360 // mark all the nodes in the dag as barren (true)
362
363 // mark all the ancestors of the evidence and targets as non-barren
364 List< NodeId > nodes_to_examine;
365 int nb_non_barren = 0;
366 for (const auto node: *_observed_nodes_)
367 nodes_to_examine.insert(node);
368 for (const auto node: *_target_nodes_)
369 nodes_to_examine.insert(node);
370
371 while (!nodes_to_examine.empty()) {
372 const NodeId node = nodes_to_examine.front();
373 nodes_to_examine.popFront();
374 if (barren_mark[node]) {
375 barren_mark[node] = false;
376 ++nb_non_barren;
377 for (const auto par: _dag_->parents(node))
378 nodes_to_examine.insert(par);
379 }
380 }
381
382 // here, all the nodes marked true are barren
383 NodeSet barren_nodes(_dag_->sizeNodes() - nb_non_barren);
384 for (const auto& marked_pair: barren_mark)
385 if (marked_pair.second) barren_nodes.insert(marked_pair.first);
386
387 return barren_nodes;
388 }
389
390} /* namespace gum */
Detect barren nodes for inference in Bayesian networks.
const NodeSet & parents(NodeId id) const
returns the set of nodes with arc ingoing to a given node
NodeSet children(const NodeSet &ids) const
returns the set of nodes which consists in the node and its parents returns the set of children of a ...
virtual void eraseArc(const Arc &arc)
removes an arc from the ArcGraphPart
The base class for all directed edges.
GUM_NODISCARD NodeId head() const
returns the head of the arc
GUM_NODISCARD NodeId first() const
returns one extremal node ID (whichever one it is is unspecified)
GUM_NODISCARD NodeId tail() const
returns the tail of the arc
Arc(NodeId tail, NodeId head)
basic constructor. Creates tail -> head.
const NodeSet * _observed_nodes_
the set of observed nodes
const DAG * _dag_
the DAG on which we compute the barren nodes
BarrenNodesFinder & operator=(const BarrenNodesFinder &from)
copy operator
NodeSet barrenNodes()
returns the set of barren nodes
BarrenNodesFinder(const DAG *dag)
default constructor
const NodeSet * _target_nodes_
the set of targeted nodes
Basic graph of cliques.
Definition cliqueGraph.h:77
const NodeSet & separator(const Edge &edge) const
returns the separator included in a given edge
const NodeSet & clique(const NodeId idClique) const
returns the set of nodes included into a given clique
Base class for dag.
Definition DAG.h:121
void addArc(NodeId tail, NodeId head) final
insert a new arc into the directed graph
Definition DAG_inl.h:75
void eraseNode(const NodeId id) override
remove a node and its adjacent arcs from the graph
Definition diGraph_inl.h:93
const EdgeSet & edges() const
returns the set of edges stored within the EdgeGraphPart
Edge(NodeId aN1, NodeId aN2)
constructs a new edge (aN1,aN2)
value_type & insert(const Key &key, const Val &val)
Adds a new element (actually a copy of this element) into the hash table.
Generic doubly linked lists.
Definition list.h:378
Val & front() const
Returns a reference to first element of a list, if any.
Definition list_tpl.h:1694
Val & insert(const Val &val)
Inserts a new element at the end of the chained list (alias of pushBack).
Definition list_tpl.h:1508
bool empty() const noexcept
Returns a boolean indicating whether the chained list is empty.
Definition list_tpl.h:1822
void popFront()
Removes the first element of a List, if any.
Definition list_tpl.h:1816
Size size() const
alias for sizeNodes
Size sizeNodes() const
returns the number of nodes in the NodeGraphPart
NodeProperty< VAL > nodesPropertyFromVal(const VAL &a, Size size=0) const
a method to create a hashMap with key:NodeId and value:VAL
bool exists(const Key &k) const
Indicates whether a given elements belong to the set.
Definition set_tpl.h:504
iterator begin() const
The usual unsafe begin iterator to parse the set.
Definition set_tpl.h:409
void insert(const Key &k)
Inserts a new element into the set.
Definition set_tpl.h:510
iterator_safe beginSafe() const
The usual safe begin iterator to parse the set.
Definition set_tpl.h:385
void erase(const Key &k)
Erases an element from the set.
Definition set_tpl.h:553
Size size() const noexcept
Returns the number of elements in the set.
Definition set_tpl.h:607
static const iterator_safe & endSafe() noexcept
The usual safe end iterator to parse the set.
Definition set_tpl.h:397
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Definition types.h:74
Size Idx
Type for indexes.
Definition types.h:79
Size NodeId
Type for node ids.
HashTable< Arc, VAL > ArcProperty
Property on graph elements.
Set< Arc > ArcSet
Some typdefs and define for shortcuts ...
HashTable< NodeId, VAL > NodeProperty
Property on graph elements.
Set< NodeId > NodeSet
Some typdefs and define for shortcuts ...
gum is the global namespace for all aGrUM entities
Definition agrum.h:46