aGrUM 3.2.0
a C++ library for (probabilistic) graphical models
KTBNLearner.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
78
79#ifndef GUM_LEARNING_KTBN_LEARNER_H
80#define GUM_LEARNING_KTBN_LEARNER_H
81
82#include <memory>
83#include <string>
84#include <utility>
85#include <vector>
86
87#include <agrum/agrum.h>
88
91
92#include <string_view>
93#include <unordered_set>
94
95namespace gum {
96
97 namespace learning {
98
114 template < GUM_Numeric GUM_SCALAR >
115 class KTBNLearner: public IKTBNLearner< GUM_SCALAR > {
116 public:
117 // #######################################################################
119 // #######################################################################
121
156 KTBNLearner(std::string_view dirPath,
157 std::string_view csvBaseName,
159 Size k,
160 const std::unordered_set< std::string >& atemporalVars,
161 const std::vector< std::string >& missingSymbols = {"?"},
162 bool induceTypes = true,
163 bool ignoreMissingSymbols = false);
164
207 KTBNLearner(std::string_view dirPath,
208 std::string_view csvBaseName,
210 Size k,
211 const std::vector< std::string >& missingSymbols = {"?"},
212 bool induceTypes = true,
213 bool ignoreMissingSymbols = false);
214
241 KTBNLearner(std::string_view dirPath,
242 std::string_view csvBaseName,
244 Size k,
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);
249
251 ~KTBNLearner();
252
254 // #######################################################################
256 // #######################################################################
258
260 KTBN< GUM_SCALAR > learnKTBN() override;
261
265 KTBN< GUM_SCALAR > learnParameters(const KTBN< GUM_SCALAR >& structure,
266 bool takeIntoAccountScore = true);
267
269 // #######################################################################
271 // #######################################################################
273
280 void useScorefNML() override;
281
284 std::string checkScorePriorCompatibility() const;
285
287 // #######################################################################
289 // #######################################################################
291
295 Size nb_decrease = 2) override;
297
299 // #######################################################################
301 // #######################################################################
303
307
312 std::vector< std::pair< std::string, std::string > > latentVariables() const;
313
315 // #######################################################################
317 // #######################################################################
319
320 KTBNLearner< GUM_SCALAR >& useSmoothingPrior(double weight = 1.0) override;
321
323 // #######################################################################
325 // #######################################################################
327
330 KTBNLearner< GUM_SCALAR >& addForbiddenArc(std::string_view tailNode,
331 std::string_view headNode) override;
332
336 KTBNLearner< GUM_SCALAR >& addForbiddenArc(std::string_view tailBase,
337 int tailSlice,
338 std::string_view headBase,
339 int headSlice) override;
340
342 KTBNLearner< GUM_SCALAR >& eraseForbiddenArc(std::string_view tailNode,
343 std::string_view headNode) override;
344
346 KTBNLearner< GUM_SCALAR >& eraseForbiddenArc(std::string_view tailBase,
347 int tailSlice,
348 std::string_view headBase,
349 int headSlice) override;
350
352 KTBNLearner< GUM_SCALAR >& addMandatoryArc(std::string_view tailNode,
353 std::string_view headNode) override;
354
356 KTBNLearner< GUM_SCALAR >& addMandatoryArc(std::string_view tailBase,
357 int tailSlice,
358 std::string_view headBase,
359 int headSlice) override;
360
362 KTBNLearner< GUM_SCALAR >& eraseMandatoryArc(std::string_view tailNode,
363 std::string_view headNode) override;
364
366 KTBNLearner< GUM_SCALAR >& eraseMandatoryArc(std::string_view tailBase,
367 int tailSlice,
368 std::string_view headBase,
369 int headSlice) override;
370
375 KTBNLearner< GUM_SCALAR >& addForbiddenIntraSliceArc(std::string_view tailBase,
376 std::string_view headBase) override;
377
381 std::string_view headBase) override;
382
387 KTBNLearner< GUM_SCALAR >& addForbiddenArcAllSlices(std::string_view tailBase,
388 std::string_view headBase) override;
389
391 KTBNLearner< GUM_SCALAR >& eraseForbiddenArcAllSlices(std::string_view tailBase,
392 std::string_view headBase) override;
393
395 KTBNLearner< GUM_SCALAR >& addNoParentNode(std::string_view base, int slice) override;
396
398 KTBNLearner< GUM_SCALAR >& addNoParentNode(std::string_view name) override;
399
401 KTBNLearner< GUM_SCALAR >& eraseNoParentNode(std::string_view base, int slice) override;
402
404 KTBNLearner< GUM_SCALAR >& eraseNoParentNode(std::string_view name) override;
405
407 KTBNLearner< GUM_SCALAR >& addNoChildrenNode(std::string_view base, int slice) override;
408
410 KTBNLearner< GUM_SCALAR >& addNoChildrenNode(std::string_view name) override;
411
413 KTBNLearner< GUM_SCALAR >& eraseNoChildrenNode(std::string_view base, int slice) override;
414
416 KTBNLearner< GUM_SCALAR >& eraseNoChildrenNode(std::string_view name) override;
417
419 KTBNLearner< GUM_SCALAR >& addPossibleEdge(std::string_view tailBase,
420 int tailSlice,
421 std::string_view headBase,
422 int headSlice) override;
423
425 KTBNLearner< GUM_SCALAR >& addPossibleEdge(std::string_view tail,
426 std::string_view head) override;
427
429 KTBNLearner< GUM_SCALAR >& erasePossibleEdge(std::string_view tailBase,
430 int tailSlice,
431 std::string_view headBase,
432 int headSlice) override;
433
435 KTBNLearner< GUM_SCALAR >& erasePossibleEdge(std::string_view tail,
436 std::string_view head) override;
437
439 KTBNLearner< GUM_SCALAR >& allowArcAdditions(bool allow = true) override;
440
442 KTBNLearner< GUM_SCALAR >& allowArcDeletions(bool allow = true) override;
443
445 KTBNLearner< GUM_SCALAR >& allowArcReversals(bool allow = true) override;
446
448 KTBNLearner< GUM_SCALAR >& setMaxIndegree(Size max_indegree) override;
449
451 // #######################################################################
453 // #######################################################################
455
457 Size k() const;
458
461 Size nbCols() const;
462
466 std::vector< Size > nbRows() const;
467
469 bool isConstraintBased() const;
470
472 bool isScoreBased() const;
473
475 std::string toString() const;
476
478 std::vector< std::tuple< std::string, std::string, std::string > > state() const;
479
482 void copyState(const KTBNLearner< GUM_SCALAR >& learner);
483
485 // #######################################################################
487 // #######################################################################
489
492 Size nbSamples() const;
493
500 bool hasMissingValues() const;
501
548 Size nbDroppedRows() const;
549
553 bool isIgnoringMissingSymbols() const;
554
557 std::vector< std::string > names() const;
558
560 std::vector< std::size_t > domainSizes() const;
561
568 Size domainSize(std::string_view base) const;
569
571
572 private:
574 std::unique_ptr< BNLearner< GUM_SCALAR > > _transitionLearner_;
575
577 std::unique_ptr< BNLearner< GUM_SCALAR > > _initialLearner_;
578
580 std::unique_ptr< BNLearner< GUM_SCALAR > > _atemporalLearner_;
581
588
589 KTBN< GUM_SCALAR > _prior_ktbn_;
590
594
598 std::vector< Size > _nbTimeSlices_;
599
608
612 template < class F >
613 void _forEachLearner_(F&& f);
614
633 template < class F >
634 void _forOwningLearner_(std::string_view tail, std::string_view head, F&& f);
635
647 template < class F >
648 void
649 _forEachAllSlicesPair_(std::string_view tailBase, std::string_view headBase, F&& f) const;
650
651 // ----- construction (runs once, in the constructor) -----
652
660 void _build_(std::string_view dirPath,
661 std::string_view csvBaseName,
663 const std::vector< std::string >& missingSymbols);
664
673 static std::unordered_set< std::string >
674 _inferAtemporalVars_(std::string_view dirPath,
675 std::string_view csvBaseName,
677 Size k,
678 const std::vector< std::string >& missingSymbols);
679
685 static KTBN< GUM_SCALAR >
686 _buildPriorFromCSV_(std::string_view dirPath,
687 std::string_view csvBaseName,
688 Size k,
689 const std::unordered_set< std::string >& atemporalVars,
690 const std::vector< std::string >& missingSymbols,
691 bool induceTypes);
692
698 static KTBN< GUM_SCALAR >
700 const BayesNet< GUM_SCALAR >& bn,
701 const std::unordered_set< std::string >& atemporalVars);
702
703 // ----- name encoding -----
704
707 const std::unordered_set< std::string >& _atemporalVarNames_() const override;
708
711 bool _isKnownBase_(std::string_view base) const override;
712
715 using IKTBNLearner< GUM_SCALAR >::_encode_;
716 using IKTBNLearner< GUM_SCALAR >::_determineNode_;
718 using IKTBNLearner< GUM_SCALAR >::_checkBaseIsTemporal_;
719
721 KTBN< GUM_SCALAR > _assemble_(const BayesNet< GUM_SCALAR >& transitionBN,
722 const BayesNet< GUM_SCALAR >& initialBN,
723 const BayesNet< GUM_SCALAR >& atemporalBN) const;
724
725 // forbidden copies / moves
730 };
731
732#ifndef GUM_NO_EXTERN_TEMPLATE_CLASS
733 extern template class GUM_PUBLIC_KTBN KTBNLearner< double >;
734#endif
735
736 } /* namespace learning */
737} /* namespace gum */
738
740
741#endif /* GUM_LEARNING_KTBN_LEARNER_H */
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.
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.
Definition types.h:74
include the inlined functions if necessary
Definition CSVParser.h:55
template class GUM_PUBLIC_KTBN KTBNLearner< double >
gum is the global namespace for all aGrUM entities
Definition agrum.h:46