79#ifndef GUM_LEARNING_KTBN_LEARNER_H
80#define GUM_LEARNING_KTBN_LEARNER_H
93#include <unordered_set>
114 template < GUM_Numeric GUM_SCALAR >
157 std::string_view csvBaseName,
160 const std::unordered_set< std::string >& atemporalVars,
161 const std::vector< std::string >& missingSymbols = {
"?"},
162 bool induceTypes =
true,
163 bool ignoreMissingSymbols =
false);
208 std::string_view csvBaseName,
211 const std::vector< std::string >& missingSymbols = {
"?"},
212 bool induceTypes =
true,
213 bool ignoreMissingSymbols =
false);
242 std::string_view csvBaseName,
245 const BayesNet< GUM_SCALAR >& bn,
246 const std::unordered_set< std::string >& atemporalVars = {},
247 const std::vector< std::string >& missingSymbols = {
"?"},
248 bool ignoreMissingSymbols =
false);
265 KTBN< GUM_SCALAR >
learnParameters(
const KTBN< GUM_SCALAR >& structure,
266 bool takeIntoAccountScore =
true);
295 Size nb_decrease = 2)
override;
312 std::vector< std::pair< std::string, std::string > >
latentVariables()
const;
331 std::string_view headNode)
override;
338 std::string_view headBase,
339 int headSlice)
override;
343 std::string_view headNode)
override;
348 std::string_view headBase,
349 int headSlice)
override;
353 std::string_view headNode)
override;
358 std::string_view headBase,
359 int headSlice)
override;
363 std::string_view headNode)
override;
368 std::string_view headBase,
369 int headSlice)
override;
376 std::string_view headBase)
override;
381 std::string_view headBase)
override;
388 std::string_view headBase)
override;
392 std::string_view headBase)
override;
421 std::string_view headBase,
422 int headSlice)
override;
426 std::string_view head)
override;
431 std::string_view headBase,
432 int headSlice)
override;
436 std::string_view head)
override;
466 std::vector< Size >
nbRows()
const;
478 std::vector< std::tuple< std::string, std::string, std::string > >
state()
const;
557 std::vector< std::string >
names()
const;
660 void _build_(std::string_view dirPath,
661 std::string_view csvBaseName,
663 const std::vector< std::string >& missingSymbols);
673 static std::unordered_set< std::string >
675 std::string_view csvBaseName,
678 const std::vector< std::string >& missingSymbols);
685 static KTBN< GUM_SCALAR >
687 std::string_view csvBaseName,
689 const std::unordered_set< std::string >& atemporalVars,
690 const std::vector< std::string >& missingSymbols,
698 static KTBN< GUM_SCALAR >
700 const BayesNet< GUM_SCALAR >& bn,
701 const std::unordered_set< std::string >& atemporalVars);
721 KTBN< GUM_SCALAR >
_assemble_(
const BayesNet< GUM_SCALAR >& transitionBN,
722 const BayesNet< GUM_SCALAR >& initialBN,
723 const BayesNet< GUM_SCALAR >& atemporalBN)
const;
732#ifndef GUM_NO_EXTERN_TEMPLATE_CLASS
A basic pack of learning algorithms that can easily be used.
Common configuration interface for k-TBN learners.
Implementation of the KTBNLearner class.
Pure-virtual configuration interface shared by all k-TBN learners.
void _checkBaseIsTemporal_(std::string_view base, std::string_view context) const
Throw InvalidArgument unless base is a known temporal base. context completes "cannot appear in <cont...
void _checkArcTemporallyFeasible_(std::string_view tail, std::string_view head, std::string_view action) const
Reject an arc the k-TBN definition can never contain, so eraseForbiddenArc and addMandatoryArc both f...
std::string _encode_(std::string_view base, int slice) const
(base, slice) -> engine name ("A[1]" / atemporal engine name). Pure function, shared by every learner...
std::pair< std::string, int > _determineNode_(const std::string &name) const
engine name -> (base, slice); atemporal names map to KTBN::ATEMPORAL. Shared by every learner; only t...
static KTBN< GUM_SCALAR > _buildPriorFromCSV_(std::string_view dirPath, std::string_view csvBaseName, Size k, const std::unordered_set< std::string > &atemporalVars, const std::vector< std::string > &missingSymbols, bool induceTypes)
called in the member-initialiser list of the k-CSV constructor: opens the first trajectory CSV,...
std::vector< std::string > names() const
Base names (no slice suffix), one entry per base variable (temporal or atemporal),...
void copyState(const KTBNLearner< GUM_SCALAR > &learner)
Copy all score/algorithm/prior/constraint settings from another KTBNLearner (does not copy the databa...
KTBNLearner< GUM_SCALAR > & useNMLCorrection() override
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
const std::unordered_set< std::string > & _atemporalVarNames_() const override
atemporal base names for IKTBNLearner's shared encode/_determineNode_; read straight from the prior k...
std::unique_ptr< BNLearner< GUM_SCALAR > > _initialLearner_
learns the initial slices 0..k-2
KTBNLearner< GUM_SCALAR > & eraseForbiddenArcAllSlices(std::string_view tailBase, std::string_view headBase) override
Undo a previous addForbiddenArcAllSlices.
void _build_(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, const std::vector< std::string > &missingSymbols)
reads every trajectory, builds the three DatabaseTables (sliding window, initial-slice flattening,...
KTBNLearner< GUM_SCALAR > & erasePossibleEdge(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice) override
Undo a previous addPossibleEdge.
KTBNLearner(KTBNLearner< GUM_SCALAR > &&)=delete
KTBN< GUM_SCALAR > learnKTBN() override
Full learning (structure + CPTs). Mirrors BNLearner::learnBN().
bool _ignoreMissingSymbols_
prior k-TBN: the single source of truth for k, variable domains, temporal/atemporal classification an...
std::vector< std::pair< std::string, std::string > > latentVariables() const
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
KTBNLearner< GUM_SCALAR > & addForbiddenArc(std::string_view tailNode, std::string_view headNode) override
Forbid tailNode from ever parenting headNode (engine names, e.g. "X[1]", "C").
Size nbDroppedRows() const
Number of rows dropped from the internal databases because they carried a missing symbol.
KTBNLearner< GUM_SCALAR > & useScoreAIC() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & useGreedyHillClimbing() override
static KTBN< GUM_SCALAR > _buildPriorFromBN_(Size k, const BayesNet< GUM_SCALAR > &bn, const std::unordered_set< std::string > &atemporalVars)
called in the member-initialiser list of the BN constructor: builds and returns a KTBN whose variable...
Size nbSamples() const
Number of trajectory CSV files loaded (the constructor's nbSamples).
KTBNLearner< GUM_SCALAR > & useScoreLog2Likelihood() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
std::vector< Size > nbRows() const
Number of time steps in each trajectory CSV (one entry per sample, in load order)....
bool isConstraintBased() const
True if the current structure-learning algorithm is constraint-based (e.g. MIIC).
void _forOwningLearner_(std::string_view tail, std::string_view head, F &&f)
Apply f to the ONE internal learner that can learn the arc tail -> head, chosen by its head: a slice-...
KTBN< GUM_SCALAR > _assemble_(const BayesNet< GUM_SCALAR > &transitionBN, const BayesNet< GUM_SCALAR > &initialBN, const BayesNet< GUM_SCALAR > &atemporalBN) const
glues the three parameter-learned BNs into a single k-TBN
Size nbCols() const
Number of columns in each CSV, i.e. of base variables (temporal + atemporal).
void _forEachLearner_(F &&f)
Apply f to each present internal learner (the atemporal one only when it exists). Factors out the fan...
Size _nbTemporalPossibleEdges_
counts of currently-active possible edges, split by kind: edges with at least one temporal endpoint,...
KTBNLearner< GUM_SCALAR > & addForbiddenArcAllSlices(std::string_view tailBase, std::string_view headBase) override
Forbid tailBase -> headBase at every causally-possible slice pair (every lag, not just matching slice...
std::unique_ptr< BNLearner< GUM_SCALAR > > _transitionLearner_
learns the transition kernel (arcs arriving at slice k-1)
KTBNLearner< GUM_SCALAR > & useScoreBD() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & allowArcReversals(bool allow=true) override
Allow or forbid arc reversals during structure search.
KTBNLearner< GUM_SCALAR > & addNoChildrenNode(std::string_view base, int slice) override
Declare a single (base, slice) node as a leaf (no children).
std::vector< std::size_t > domainSizes() const
Domain sizes of the base variables, in the same column order as names().
KTBNLearner(const KTBNLearner< GUM_SCALAR > &)=delete
KTBN< GUM_SCALAR > learnParameters(const KTBN< GUM_SCALAR > &structure, bool takeIntoAccountScore=true)
CPTs only, using the arc structure of structure. structure must have the same base variables (names a...
Size _nbDroppedRows_
number of time steps (rows) in each trajectory CSV, in load order. rows build() dropped because they ...
KTBN< GUM_SCALAR > _prior_ktbn_
KTBNLearner< GUM_SCALAR > & allowArcDeletions(bool allow=true) override
Allow or forbid arc deletions during structure search.
KTBNLearner< GUM_SCALAR > & useExtendedGreedyHillClimbing() override
std::string checkScorePriorCompatibility() const
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & useNoCorrection() override
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
bool isScoreBased() const
True if the current structure-learning algorithm is score-based (e.g. BIC, AIC).
KTBNLearner< GUM_SCALAR > & allowArcAdditions(bool allow=true) override
Allow or forbid arc additions during structure search.
Size domainSize(std::string_view base) const
Domain size of the base variable base (e.g. "X", "C"). Engine names (e.g. "X[1]") are also accepted.
KTBNLearner< GUM_SCALAR > & operator=(KTBNLearner< GUM_SCALAR > &&)=delete
KTBNLearner< GUM_SCALAR > & eraseForbiddenArc(std::string_view tailNode, std::string_view headNode) override
Undo a previous addForbiddenArc (engine names, e.g. "X[1]", "C").
KTBNLearner< GUM_SCALAR > & useScoreBIC() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
bool isIgnoringMissingSymbols() const
Whether build() drops the rows carrying a missing symbol.
KTBNLearner< GUM_SCALAR > & useMIIC() override
void useScorefNML() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & useScoreMDL() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & eraseNoChildrenNode(std::string_view base, int slice) override
Undo a previous addNoChildrenNode for a single (base, slice) node.
static std::unordered_set< std::string > _inferAtemporalVars_(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, Size k, const std::vector< std::string > &missingSymbols)
checks k >= 2, then delegates the actual scan to the shared IKTBNLearner::scanConstantColumns() (also...
KTBNLearner< GUM_SCALAR > & eraseForbiddenIntraSliceArc(std::string_view tailBase, std::string_view headBase) override
Undo a previous addForbiddenIntraSliceArc.
std::string toString() const
Human-readable summary of the learner's current configuration.
std::vector< Size > _nbTimeSlices_
Captured once by build() and exposed by nbRows(). This is the raw trajectory length,...
std::vector< std::tuple< std::string, std::string, std::string > > state() const
Settings as a vector of (key, value, comment) tuples (mirrors BNLearner::state()).
KTBNLearner< GUM_SCALAR > & addMandatoryArc(std::string_view tailNode, std::string_view headNode) override
Force tailNode to be a parent of headNode (engine names, e.g. "X[1]", "C").
bool hasMissingValues() const
True if any internal database contains missing values.
KTBNLearner< GUM_SCALAR > & setMaxIndegree(Size max_indegree) override
Cap the number of parents of any single node.
KTBNLearner< GUM_SCALAR > & addForbiddenIntraSliceArc(std::string_view tailBase, std::string_view headBase) override
Forbid tailBase -> headBase at every intra-slice position (i.e. tailBase[t] -> headBase[t] for all t ...
KTBNLearner< GUM_SCALAR > & addPossibleEdge(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice) override
Add a candidate edge for MIIC (only edges explicitly listed are explored).
bool _isKnownBase_(std::string_view base) const override
whether base is one of this learner's variables; read straight from the prior k-TBN,...
KTBNLearner< GUM_SCALAR > & useMDLCorrection() override
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
KTBNLearner< GUM_SCALAR > & useScoreBDeu() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
Size _nbAtemporalPossibleEdges_
KTBNLearner< GUM_SCALAR > & useSmoothingPrior(double weight=1.0) override
KTBNLearner< GUM_SCALAR > & eraseMandatoryArc(std::string_view tailNode, std::string_view headNode) override
Undo a previous addMandatoryArc (engine names, e.g. "X[1]", "C").
KTBNLearner< GUM_SCALAR > & addNoParentNode(std::string_view base, int slice) override
Declare a single (base, slice) node as a root (no parents).
KTBNLearner< GUM_SCALAR > & useLocalSearchWithTabuList(Size tabu_size=100, Size nb_decrease=2) override
KTBNLearner< GUM_SCALAR > & operator=(const KTBNLearner< GUM_SCALAR > &)=delete
std::unique_ptr< BNLearner< GUM_SCALAR > > _atemporalLearner_
learns the atemporal variables (arcs atemporal -> atemporal)
KTBNLearner(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, Size k, const std::unordered_set< std::string > &atemporalVars, const std::vector< std::string > &missingSymbols={"?"}, bool induceTypes=true, bool ignoreMissingSymbols=false)
Structure-learning constructor — variable roles supplied explicitly.
Size k() const
Order of the k-TBN being learned.
KTBNLearner< GUM_SCALAR > & eraseNoParentNode(std::string_view base, int slice) override
Undo a previous addNoParentNode for a single (base, slice) node.
void _forEachAllSlicesPair_(std::string_view tailBase, std::string_view headBase, F &&f) const
Apply f(tailSlice, headSlice) to every causally-possible slice pair of an all-slices constraint betwe...
std::size_t Size
In aGrUM, hashed values are unsigned long int.
include the inlined functions if necessary
template class GUM_PUBLIC_KTBN KTBNLearner< double >
gum is the global namespace for all aGrUM entities