aGrUM 3.2.0
a C++ library for (probabilistic) graphical models
KTBNInference.h
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
93#ifndef GUM_KTBN_INFERENCE_H
94#define GUM_KTBN_INFERENCE_H
95
96#include <map>
97#include <memory>
98#include <set>
99#include <string>
100#include <variant>
101#include <vector>
102
103#include <agrum/agrum.h>
104
107#include <agrum/KTBN/KTBN.h>
108
109#include <unordered_map>
110#include <unordered_set>
111
112namespace gum {
113
146 template < GUM_Numeric GUM_SCALAR >
148 public:
151
152 // ===========================================================================
154 // ===========================================================================
156
162 explicit KTBNInference(const KTBN< GUM_SCALAR >* ktbn);
163
165 ~KTBNInference() = default;
166
170
172
175 using NodeKey = std::variant< std::string, std::pair< std::string, int > >;
176
177 // ===========================================================================
179 // ===========================================================================
181
198 void addIntervention(std::string_view base, int slice, const KTBNModality& value);
200 void addIntervention(std::string_view node_name, const KTBNModality& value);
201
222 void addIntervention(const std::vector< std::pair< NodeKey, KTBNModality > >& interventions);
223
225 void eraseIntervention(std::string_view base, int slice);
227 void eraseIntervention(std::string_view node_name);
229 void clearInterventions();
230
232 bool hasIntervention(std::string_view base, int slice) const;
234 bool hasIntervention(std::string_view node_name) const;
235
237 // ===========================================================================
239 // ===========================================================================
241
259 void addObservation(std::string_view base, int slice, const KTBNModality& value);
261 void addObservation(std::string_view node_name, const KTBNModality& value);
262
274 void addObservation(std::string_view base,
275 int slice,
276 const std::vector< GUM_SCALAR >& likelihood);
278 void addObservation(std::string_view node_name, const std::vector< GUM_SCALAR >& likelihood);
279
285 void addObservations(const std::vector< std::pair< NodeKey, KTBNModality > >& observations);
286
288 void eraseObservation(std::string_view base, int slice);
290 void eraseObservation(std::string_view node_name);
292 void clearObservation();
293
295 bool hasObservation(std::string_view base, int slice) const;
297 bool hasObservation(std::string_view node_name) const;
301 bool hasObservation() const;
302
304 // ===========================================================================
306 // ===========================================================================
308
319 void addTarget(std::string_view base);
320
323 void eraseTarget(std::string_view base);
325 void clearTargets();
326
328 bool isTarget(std::string_view base) const;
329
333 bool isInTargetMode() const;
334
336 // ===========================================================================
338 // ===========================================================================
340
361 void makeInference(Size nbTimeSlices);
362
374 const Tensor< GUM_SCALAR >& posterior(std::string_view base, int slice);
376 const Tensor< GUM_SCALAR >& posterior(std::string_view node_name);
377
388 const std::vector< Tensor< GUM_SCALAR > >& posteriors(std::string_view base);
389
394 GUM_SCALAR logObservationProbability();
395
399 GUM_SCALAR observationProbability();
400
402 // ===========================================================================
404 // ===========================================================================
406
408 const KTBN< GUM_SCALAR >& ktbn() const;
409
411 std::string toString() const;
412
416 const JunctionTree& windowJunctionTree() const;
417
420 Size interfaceSize() const;
421
423
424 private:
429 struct _Series_ {
430 std::vector< std::unique_ptr< DiscreteVariable > > vars;
431 std::vector< Tensor< GUM_SCALAR > > tensors;
432 };
433
437 struct _Slot_ {
438 int base;
439 int lag;
440 bool operator==(const _Slot_& o) const;
441 };
442
447 struct _Window_ {
449 std::vector< _Slot_ > slotOfNode;
457 std::vector< NodeId > bfs;
458 std::unordered_map< NodeId, NodeId > parentOf;
460 std::unordered_map< NodeId, std::vector< int > > factorsOf;
463 std::unordered_map< int, NodeId > selfClique;
465 std::vector< _Slot_ > Iprev, Icur;
466 };
467
469 const KTBN< GUM_SCALAR >* _ktbn_;
470
472 int _k_;
473
475 std::map< std::string, Idx > _interventions_;
476
479 std::map< std::string, std::vector< GUM_SCALAR > > _observations_;
480
482 std::set< std::string > _targets_;
483
485 bool _targeted_mode_{false};
486
489
491 bool _done_{false};
492
494 GUM_SCALAR _logObservation_{0};
495
497 std::unordered_map< std::string, _Series_ > _posteriors_;
498
501 std::vector< std::string > _temporalSorted_;
502 std::vector< std::string > _atemporalSorted_;
503
506 std::vector< std::string > _baseNames_;
507 std::size_t _nbTemporal_{0};
509 std::unordered_map< std::string, int > _baseIdx_;
510
514 std::vector< int > _maxLag_;
515
518 std::vector< _Window_ > _windows_;
519
523 std::vector< bool > _requisite_;
524
530 mutable std::vector< std::unordered_map< NodeId, Tensor< GUM_SCALAR > > > _psiCache_;
531 mutable std::vector< bool > _psiCached_;
532
537 mutable std::unordered_set< int > _observationSlices_;
538
544 mutable std::unordered_set< int > _interventionSlices_;
545
547 mutable std::unordered_map< NodeId, Tensor< GUM_SCALAR > > _psiScratch_;
548
551 Size _psiKey_(int t) const;
552
565 mutable std::map< std::pair< std::string, int >, Tensor< GUM_SCALAR > > _kernelCache_;
566
569
571 bool _isTemporal_(const std::string& base) const;
573 bool _isAtemporal_(const std::string& base) const;
574
576 static std::string _encode_(const std::string& base, int slice);
577
581 std::pair< std::string, int > _determineNode_(const std::string& name) const;
582
584 void _validateNode_(const std::string& base, int slice) const;
585
587 const DiscreteVariable& _templateVar_(const std::string& base, int slice) const;
588
593 const DiscreteVariable* _varOfSlot_(const _Slot_& s, int t) const;
594
603 const Tensor< GUM_SCALAR >& _buildKernel_(const std::string& p, int t) const;
604
609 std::vector< _Slot_ > _familySlots_(int baseIdx, int t) const;
610
616 std::vector< _Slot_ > _interfaceAfter_(int t) const;
617
620 int _lastConsumerSlice_(int baseIdx, int s) const;
621
626 void _buildWindows_();
627
629 _Window_ _compileWindow_(const std::vector< _Slot_ >& Iprev,
630 const std::vector< _Slot_ >& Icur,
631 int t,
632 bool withAtemporalFamilies) const;
633
636 void _markRequisite_();
637
641
644 const _Window_& _windowAt_(int t) const;
645
666 const std::unordered_map< NodeId, Tensor< GUM_SCALAR > >& _windowPotentials_(const _Window_& w,
667 int t) const;
668
672 int t,
673 std::unordered_map< NodeId, Tensor< GUM_SCALAR > >& psi) const;
674
678 void _fillWindow_(const _Window_& w,
679 int t,
680 std::unordered_map< NodeId, Tensor< GUM_SCALAR > >& psi,
681 bool withTemporalEvidence = true) const;
682
689 void _propagate_(const _Window_& w,
690 int t,
691 const std::unordered_map< NodeId, Tensor< GUM_SCALAR > >& psi,
692 const Tensor< GUM_SCALAR >* inPrev,
693 const Tensor< GUM_SCALAR >* inNext,
694 bool distribute,
695 std::map< std::pair< NodeId, NodeId >, Tensor< GUM_SCALAR > >& msgs) const;
696
699 Tensor< GUM_SCALAR >
700 _belief_(const _Window_& w,
701 const std::unordered_map< NodeId, Tensor< GUM_SCALAR > >& psi,
702 const std::map< std::pair< NodeId, NodeId >, Tensor< GUM_SCALAR > >& msgs,
703 const Tensor< GUM_SCALAR >* inPrev,
704 const Tensor< GUM_SCALAR >* inNext,
705 NodeId c,
706 NodeId skipNeighbour) const;
707
713 void _snapshot_(const std::string& base, int slice, const Tensor< GUM_SCALAR >& marginal);
714
717 const _Series_& _series_(const std::string& base);
718
720 };
721
722#ifndef GUM_NO_EXTERN_TEMPLATE_CLASS
723 extern template class GUM_PUBLIC_KTBN KTBNInference< double >;
724#endif
725
726} // namespace gum
727
729
730#endif /* GUM_KTBN_INFERENCE_H */
Template implementation of gum::KTBNInference (interface algorithm).
Class representing k-order dynamic Bayesian networks (k-DBN).
Base class for discrete random variable.
void _fillWindow_(const _Window_ &w, int t, std::unordered_map< NodeId, Tensor< GUM_SCALAR > > &psi, bool withTemporalEvidence=true) const
void _snapshot_(const std::string &base, int slice, const Tensor< GUM_SCALAR > &marginal)
Snapshots marginal onto an owned, stably-named descriptor and appends it to that base's series (index...
void makeInference(Size nbTimeSlices)
Runs the interface algorithm over nbTimeSlices slices ( ) and caches, for every targeted base,...
void addObservation(std::string_view base, int slice, const KTBNModality &value)
Records a hard observation .
KTBNInference(const KTBN< GUM_SCALAR > *ktbn)
Constructor.
void _markRequisite_()
Marks the requisite bases of the current run (targets, observed bases and all their ancestors) into r...
void _applyTemporalObservations_(const _Window_ &w, int t, std::unordered_map< NodeId, Tensor< GUM_SCALAR > > &psi) const
Multiplies slice t's temporal observation likelihoods into an already-built base. Atemporal ones are ...
void clearObservation()
Removes all recorded observations.
KTBNInference(const KTBNInference< GUM_SCALAR > &)=delete
Copy is disabled (owns per-run variable descriptors and cached tensors).
std::vector< bool > _psiCached_
std::unordered_map< NodeId, Tensor< GUM_SCALAR > > _psiScratch_
Potentials of an evidence-carrying slice, rebuilt on each visit.
void eraseIntervention(std::string_view base, int slice)
Removes a recorded intervention (silent no-op if absent).
bool isInTargetMode() const
Size _horizon_
Horizon (nbTimeSlices) of the last/next run; 0 <=> makeInference never run.
const DiscreteVariable & _templateVar_(const std::string &base, int slice) const
A representative template variable of base for domain/cloning.
std::vector< bool > _requisite_
Bases actually folded by the current run: the targets, the observed nodes and all their ancestors....
void clearTargets()
Removes all targets (restores default-all-targets mode).
void addTarget(std::string_view base)
Declares a target: a base variable whose marginals we want.
GUM_SCALAR _logObservation_
log P(observation | do) of the last run.
std::unordered_set< int > _observationSlices_
Slices carrying a temporal observation. Their potentials are the periodic ones times that slice's lik...
GUM_SCALAR logObservationProbability()
for the last run: the likelihood of the observations under the (possibly mutilated) model....
_Window_ _compileWindow_(const std::vector< _Slot_ > &Iprev, const std::vector< _Slot_ > &Icur, int t, bool withAtemporalFamilies) const
Compiles one window over the given slot set / interfaces.
std::set< std::string > _targets_
Recorded targets (base names). Empty <=> default-all-targets mode.
Tensor< GUM_SCALAR > _belief_(const _Window_ &w, const std::unordered_map< NodeId, Tensor< GUM_SCALAR > > &psi, const std::map< std::pair< NodeId, NodeId >, Tensor< GUM_SCALAR > > &msgs, const Tensor< GUM_SCALAR > *inPrev, const Tensor< GUM_SCALAR > *inNext, NodeId c, NodeId skipNeighbour) const
The belief of clique c: its potential times every message reaching it, interface messages included.
std::vector< int > _maxLag_
maxLag[i]: largest lag at which the transition kernel still consumes temporal base i – how long an oc...
std::variant< std::string, std::pair< std::string, int > > NodeKey
A node designated either by its engine name ("X[2]", "C") or by its (base, slice) identity.
std::vector< std::string > _baseNames_
All bases: temporal first (indices 0.._nbTemporal_-1), then atemporal. Window slots index into this.
bool _targeted_mode_
Whether at least one explicit target has been declared.
std::unordered_map< std::string, _Series_ > _posteriors_
Cached marginal series of the last run, keyed by base name.
const JunctionTree & windowJunctionTree() const
The junction tree of the repeating window – the one compiled from the k-slice template and re-entered...
void addObservations(const std::vector< std::pair< NodeKey, KTBNModality > > &observations)
Records several observations in one call, all-or-nothing.
const std::unordered_map< NodeId, Tensor< GUM_SCALAR > > & _windowPotentials_(const _Window_ &w, int t) const
Clique potentials of the window at slice t: every requisite family's CPT (or, under an intervention,...
void _validateNode_(const std::string &base, int slice) const
Validates that (base, slice) denotes a legal node (future slices ok).
void clearInterventions()
Removes all recorded interventions.
bool hasIntervention(std::string_view base, int slice) const
const KTBN< GUM_SCALAR > & ktbn() const
Size interfaceSize() const
Size of the forward interface of the repeating window: how many node occurrences have to cross each s...
std::string toString() const
std::vector< std::unordered_map< NodeId, Tensor< GUM_SCALAR > > > _psiCache_
Memoized clique potentials for the slices that carry no temporal evidence, indexed by psiKey(t)....
static std::string _encode_(const std::string &base, int slice)
Encodes (base, slice) -> engine name (base[slice] or bare base).
std::map< std::string, Idx > _interventions_
Recorded interventions, keyed by engine name -> forced value.
KTBNInference< GUM_SCALAR > & operator=(const KTBNInference< GUM_SCALAR > &)=delete
Constructor.
const Tensor< GUM_SCALAR > & _buildKernel_(const std::string &p, int t) const
Transition-kernel tensor of process p at slice t ( ): the template kernel remapped onto the k-DBN's o...
const KTBN< GUM_SCALAR > * _ktbn_
The k-DBN (referenced, not owned).
bool isTarget(std::string_view base) const
const DiscreteVariable * _varOfSlot_(const _Slot_ &s, int t) const
The variable a window slot stands for at absolute slice t: ring slot for a temporal base,...
GUM_SCALAR observationProbability()
, i.e. exp of logObservationProbability(). Underflows to 0 on long horizons; prefer the log form.
std::vector< _Slot_ > _interfaceAfter_(int t) const
The forward interface after slice t, as slots relative to t: every requisite occurrence at a slice s...
Size _psiKey_(int t) const
Cache slot for slice t: the initial slices keep their own, the repeating window contributes one per p...
void eraseObservation(std::string_view base, int slice)
Removes a recorded observation (silent no-op if absent).
void _propagate_(const _Window_ &w, int t, const std::unordered_map< NodeId, Tensor< GUM_SCALAR > > &psi, const Tensor< GUM_SCALAR > *inPrev, const Tensor< GUM_SCALAR > *inNext, bool distribute, std::map< std::pair< NodeId, NodeId >, Tensor< GUM_SCALAR > > &msgs) const
Shafer-Shenoy pass over a filled window. inPrev / inNext are the interface messages arriving at rootD...
const _Series_ & _series_(const std::string &base)
The cached series of a targeted base, running makeInference() lazily (with the last horizon) if out o...
std::vector< _Window_ > _windows_
The compiled windows: index t for t <= k-2 (initial), index k-1 for the repeating window,...
std::size_t _nbTemporal_
bool _done_
Whether the cached posteriors are up to date.
void addIntervention(std::string_view base, int slice, const KTBNModality &value)
Records a hard intervention .
~KTBNInference()=default
Destructor.
void eraseTarget(std::string_view base)
Removes a target; when the last one is removed, default-all-targets mode is restored.
std::unordered_set< int > _interventionSlices_
Slices carrying a temporal intervention. These need a full rebuild: do(X=x) replaces the node's CPT,...
std::vector< std::string > _atemporalSorted_
int _lastConsumerSlice_(int baseIdx, int s) const
The last slice at which occurrence base[s] is still consumed (-1 if never), over both the initial fam...
void _buildWindows_()
Compiles the k window junction trees, once, from the constructor: moralise each window's families,...
std::vector< _Slot_ > _familySlots_(int baseIdx, int t) const
Parents of base at a window whose current slice is t, as slots (lag = t - parentSlice)....
std::pair< std::string, int > _determineNode_(const std::string &name) const
Cache-aware classification of an engine name -> (base, slice): a name registered as atemporal (incl....
std::map< std::pair< std::string, int >, Tensor< GUM_SCALAR > > _kernelCache_
Memoized transition kernels, keyed by (process, t % k).
int _k_
The order k, cached as int for slice arithmetic.
const std::vector< Tensor< GUM_SCALAR > > & posteriors(std::string_view base)
The whole marginal time-series of a targeted base: tensors[t] is for (a single-element vector,...
bool _isTemporal_(const std::string &base) const
std::vector< std::string > _temporalSorted_
Temporal / atemporal base names in a deterministic (sorted) order, cached once at construction (the K...
std::map< std::string, std::vector< GUM_SCALAR > > _observations_
Recorded observations, keyed by engine name -> likelihood vector (one-hot for a hard observation).
bool hasObservation(std::string_view base, int slice) const
std::unordered_map< std::string, int > _baseIdx_
name -> index into baseNames
static constexpr int ATEMPORAL
Convenience alias for the atemporal-slice sentinel.
const Tensor< GUM_SCALAR > & posterior(std::string_view base, int slice)
Returns .
bool _isAtemporal_(const std::string &base) const
const _Window_ & _windowAt_(int t) const
The window for absolute slice t: its own while t is inside the initial block, the repeating one (inde...
static constexpr int ATEMPORAL
Conventional time-slice value denoting an atemporal (static) variable.
Definition KTBN.h:200
Class for computing default triangulations of graphs.
Base class for discrete random variable.
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Definition types.h:74
Size NodeId
Type for node ids.
gum is the global namespace for all aGrUM entities
Definition agrum.h:46
CliqueGraph JunctionTree
a junction tree is a clique graph satisfying the running intersection property and such that no cliqu...
template class GUM_PUBLIC_KTBN KTBNInference< double >
A cached marginal time-series for one base: owned variable descriptors paired with their marginals,...
std::vector< std::unique_ptr< DiscreteVariable > > vars
std::vector< Tensor< GUM_SCALAR > > tensors
One node of a window template: a base (index into baseNames) at a lag behind the window's current sli...
bool operator==(const _Slot_ &o) const
A compiled window: the junction tree of , rooted at the clique holding , plus everything needed to fi...
std::unordered_map< NodeId, NodeId > parentOf
NodeId rootC
clique containing the whole outgoing interface I_t
std::unordered_map< NodeId, std::vector< int > > factorsOf
clique -> base indices whose family factor is multiplied in there
std::vector< _Slot_ > Icur
std::vector< _Slot_ > Iprev
the two interfaces, as slot lists
NodeId rootD
clique containing the whole incoming interface I_{t-1}
JunctionTree jt
the junction tree over those nodes
std::vector< _Slot_ > slotOfNode
template graph NodeId -> slot it stands for
std::vector< NodeId > bfs
cliques in BFS order from rootC, and each one's parent in that rooting
std::unordered_map< int, NodeId > selfClique
base index -> clique holding that base's own slot (lag 0 / atemporal), for reading its posterior and ...
A parent's value in gum::KTBN::fillCPT(): a modality index or a modality label.
Definition KTBN.h:95