93#ifndef GUM_KTBN_INFERENCE_H
94#define GUM_KTBN_INFERENCE_H
109#include <unordered_map>
110#include <unordered_set>
146 template < GUM_Numeric GUM_SCALAR >
175 using NodeKey = std::variant< std::string, std::pair< std::string, int > >;
222 void addIntervention(
const std::vector< std::pair< NodeKey, KTBNModality > >& interventions);
276 const std::vector< GUM_SCALAR >& likelihood);
278 void addObservation(std::string_view node_name,
const std::vector< GUM_SCALAR >& likelihood);
285 void addObservations(
const std::vector< std::pair< NodeKey, KTBNModality > >& observations);
328 bool isTarget(std::string_view base)
const;
374 const Tensor< GUM_SCALAR >&
posterior(std::string_view base,
int slice);
376 const Tensor< GUM_SCALAR >&
posterior(std::string_view node_name);
388 const std::vector< Tensor< GUM_SCALAR > >&
posteriors(std::string_view base);
408 const KTBN< GUM_SCALAR >&
ktbn()
const;
430 std::vector< std::unique_ptr< DiscreteVariable > >
vars;
457 std::vector< NodeId >
bfs;
460 std::unordered_map< NodeId, std::vector< int > >
factorsOf;
530 mutable std::vector< std::unordered_map< NodeId, Tensor< GUM_SCALAR > > >
_psiCache_;
547 mutable std::unordered_map< NodeId, Tensor< GUM_SCALAR > >
_psiScratch_;
565 mutable std::map< std::pair< std::string, int >, Tensor< GUM_SCALAR > >
_kernelCache_;
576 static std::string
_encode_(
const std::string& base,
int slice);
581 std::pair< std::string, int >
_determineNode_(
const std::string& name)
const;
603 const Tensor< GUM_SCALAR >&
_buildKernel_(
const std::string& p,
int t)
const;
609 std::vector< _Slot_ >
_familySlots_(
int baseIdx,
int t)
const;
630 const std::vector< _Slot_ >& Icur,
632 bool withAtemporalFamilies)
const;
673 std::unordered_map<
NodeId, Tensor< GUM_SCALAR > >& psi)
const;
680 std::unordered_map<
NodeId, Tensor< GUM_SCALAR > >& psi,
681 bool withTemporalEvidence =
true)
const;
691 const std::unordered_map<
NodeId, Tensor< GUM_SCALAR > >& psi,
692 const Tensor< GUM_SCALAR >* inPrev,
693 const Tensor< GUM_SCALAR >* inNext,
695 std::map< std::pair< NodeId, NodeId >, Tensor< GUM_SCALAR > >& msgs)
const;
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,
706 NodeId skipNeighbour)
const;
713 void _snapshot_(
const std::string& base,
int slice,
const Tensor< GUM_SCALAR >& marginal);
722#ifndef GUM_NO_EXTERN_TEMPLATE_CLASS
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,...
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.
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.
Size NodeId
Type for node ids.
gum is the global namespace for all aGrUM entities
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.