aGrUM 3.2.0
a C++ library for (probabilistic) graphical models
IBNLearner.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
52#ifndef GUM_LEARNING_GENERIC_BN_LEARNER_H
53#define GUM_LEARNING_GENERIC_BN_LEARNER_H
54
55#include <memory>
56#include <sstream>
57
58#include <agrum/agrum.h>
59
93
94namespace gum::learning {
96
104 class GUM_PUBLIC_BN IBNLearner:
106 public ThreadNumberManager {
107 public:
109 enum class ScoreType { AIC, BD, BDeu, BIC, fNML, K2, LOG2LIKELIHOOD, MDL };
110
113 enum class ParamEstimatorType { ML };
114
117 NO_prior,
118 SMOOTHING,
119 DIRICHLET_FROM_DATABASE,
120 DIRICHLET_FROM_BAYESNET,
121 BDEU
122 };
123
125 enum class AlgoType {
127 GREEDY_HILL_CLIMBING,
128 LOCAL_SEARCH_WITH_TABU_LIST,
129 MIIC,
132 EXTENDED_GREEDY_HILL_CLIMBING,
133 GREEDY_THICK_THINNING
134 };
135
137 static constexpr double default_EM_noise{0.1};
138
140 class GUM_PUBLIC_BN Database {
141 public:
142 // ########################################################################
144 // ########################################################################
146
148
158 explicit Database(std::string_view file,
159 const std::vector< std::string >& missing_symbols,
160 const bool induceTypes = false);
161
163
166 explicit Database(const DatabaseTable& db);
167
169
179 Database(std::string_view filename,
180 const Database& score_database,
181 const std::vector< std::string >& missing_symbols);
182
184
190 template < GUM_Numeric GUM_SCALAR >
191 Database(std::string_view filename,
193 const std::vector< std::string >& missing_symbols);
194
196 Database(const Database& from);
197
199 Database(Database&& from);
200
202 ~Database();
203
205
206 // ########################################################################
208 // ########################################################################
210
212 Database& operator=(const Database& from);
213
215 Database& operator=(Database&& from);
216
218
219 // ########################################################################
221 // ########################################################################
223
226
228 const std::vector< std::size_t >& domainSizes() const;
229
231 const std::vector< std::string >& names() const;
232
234 NodeId idFromName(std::string_view var_name) const;
235
237 const std::string& nameFromId(NodeId id) const;
238
240 const DatabaseTable& databaseTable() const;
241
244 void setDatabaseWeight(const double new_weight);
245
248
250 const std::vector< std::string >& missingSymbols() const;
251
253 std::size_t nbRows() const;
254
256 std::size_t size() const;
257
259
262 void setWeight(const std::size_t i, const double weight);
263
265
267 double weight(const std::size_t i) const;
268
270 double weight() const;
271
273
274 protected:
277
280
282 std::vector< std::size_t > _domain_sizes_;
283
286
289
292
293 private:
294 // returns the set of variables as a BN. This is convenient for
295 // the constructors of prior Databases
296 template < GUM_Numeric GUM_SCALAR >
297 BayesNet< GUM_SCALAR > _BNVars_() const;
298 };
299
300 // ##########################################################################
302 // ##########################################################################
304
317 IBNLearner(std::string_view filename,
318 const std::vector< std::string >& missingSymbols,
319 bool induceTypes = true);
320
321 explicit IBNLearner(const DatabaseTable& db);
322
340 template < GUM_Numeric GUM_SCALAR >
341 IBNLearner(std::string_view filename,
343 const std::vector< std::string >& missing_symbols);
344
346 IBNLearner(const IBNLearner&);
347
350
352 ~IBNLearner() override;
353
355
356 // ##########################################################################
358 // ##########################################################################
360
362 IBNLearner& operator=(const IBNLearner&);
363
365 IBNLearner& operator=(IBNLearner&&);
366
368
369 // ##########################################################################
371 // ##########################################################################
373
375 DAG learnDAG();
376
379 PDAG learnPDAG();
380
383 PAG learnPAG();
384
386 void setInitialDAG(const DAG&);
387
389 DAG initialDAG();
390
392 const std::vector< std::string >& names() const;
393
395 const std::vector< std::size_t >& domainSizes() const;
396 Size domainSize(NodeId var) const;
397 Size domainSize(std::string_view var) const;
398
400
404 NodeId idFromName(std::string_view var_name) const;
405
407 const DatabaseTable& database() const;
408
411 void setDatabaseWeight(const double new_weight);
412
414
417 void setRecordWeight(const std::size_t i, const double weight);
418
420
422 double recordWeight(const std::size_t i) const;
423
425 double databaseWeight() const;
426
428 const std::string& nameFromId(NodeId id) const;
429
431
437 void useDatabaseRanges(const std::vector< std::pair< std::size_t, std::size_t > >& new_ranges);
438
440 void clearDatabaseRanges();
441
443
446 const std::vector< std::pair< std::size_t, std::size_t > >& databaseRanges() const;
447
449
469 std::pair< std::size_t, std::size_t > useCrossValidationFold(const std::size_t learning_fold,
470 const std::size_t k_fold);
471
472
480 std::pair< double, double >
481 chi2(NodeId id1, NodeId id2, const std::vector< NodeId >& knowing = {});
489 std::pair< double, double > chi2(std::string_view name1,
490 std::string_view name2,
491 const std::vector< std::string >& knowing = {});
492
500 std::pair< double, double >
501 G2(NodeId id1, NodeId id2, const std::vector< NodeId >& knowing = {});
509 std::pair< double, double > G2(std::string_view name1,
510 std::string_view name2,
511 const std::vector< std::string >& knowing = {});
512
520 double logLikelihood(const std::vector< NodeId >& vars,
521 const std::vector< NodeId >& knowing = {});
522
530 double logLikelihood(const std::vector< std::string >& vars,
531 const std::vector< std::string >& knowing = {});
532
544 double mutualInformation(NodeId id1, NodeId id2, const std::vector< NodeId >& knowing = {});
545
557 double mutualInformation(std::string_view var1,
558 std::string_view var2,
559 const std::vector< std::string >& knowing = {});
560
561
574 double correctedMutualInformation(NodeId id1,
575 NodeId id2,
576 const std::vector< NodeId >& knowing = {});
577
592 double correctedMutualInformation(std::string_view var1,
593 std::string_view var2,
594 const std::vector< std::string >& knowing = {});
595
604 double score(NodeId vars, const std::vector< NodeId >& knowing = {});
605
615 double score(std::string_view vars, const std::vector< std::string >& knowing = {});
616
622 std::vector< double > rawPseudoCount(const std::vector< NodeId >& vars);
623
629 std::vector< double > rawPseudoCount(const std::vector< std::string >& vars);
634 Size nbCols() const;
635
640 Size nbRows() const;
641
659 void useEM(const double epsilon, const double noise = default_EM_noise);
660
675 void useEMWithRateCriterion(const double epsilon, const double noise = default_EM_noise);
676
690 void useEMWithDiffCriterion(const double epsilon, const double noise = default_EM_noise);
691
693 void forbidEM();
694
696 bool isUsingEM() const;
697
707 EMApproximationScheme& EM();
708
710 ApproximationSchemeSTATE EMState() const;
711
713 std::string EMStateMessage() const;
714
716 bool hasMissingValues() const;
717
719
720 // ##########################################################################
722 // ##########################################################################
724
726 void useScoreAIC();
727
729 void useScoreBD();
730
732 void useScoreBDeu();
733
735 void useScoreBIC();
736
738 void useScorefNML();
739
741 void useScoreK2();
742
744 void useScoreLog2Likelihood();
745
747 void useScoreMDL();
748
750
751 // ##########################################################################
753 // ##########################################################################
755
757 void useNoPrior();
758
760
763 void useBDeuPrior(double weight = 1.0);
764
766
769 void useSmoothingPrior(double weight = 1);
770
772 void useDirichletPrior(std::string_view filename, double weight = 1);
773
775
777 std::string checkScorePriorCompatibility() const;
779
780 // ##########################################################################
782 // ##########################################################################
784
786 void useGreedyHillClimbing();
787
789
801 void useExtendedGreedyHillClimbing();
802
804
807 void useLocalSearchWithTabuList(Size tabu_size = 100, Size nb_decrease = 2);
808
810 void useK2(const Sequence< NodeId >& order);
811
813 void useK2(const std::vector< NodeId >& order);
814
816 void useMIIC();
817
819 void usePC();
820
822 void useFCI();
823
825 void useGreedyThickThinning();
826
828 void setGreedyThickThinningReversals(bool allow);
829
831 bool greedyThickThinningReversals() const;
832
834 bool isConstraintBased() const;
835
837 bool isScoreBased() const;
838
839 // ##########################################################################
841 // ##########################################################################
843
848 void allowArcAdditions(bool allow = true);
849
858 void allowArcDeletions(bool allow = true);
859
866 void allowArcReversals(bool allow = true);
867
873 void allowArcTriangleDeletions(bool allow = true);
874
876
877 // ##########################################################################
879 // ##########################################################################
883 void useNMLCorrection();
886 void useMDLCorrection();
889 void useNoCorrection();
890
893 std::vector< Arc > latentVariables() const;
894
896
897 // ##########################################################################
899 // ##########################################################################
901
904 void useChi2Test();
905
908 void useG2Test();
909
912 void setPCAlpha(double alpha);
913
916 void setPCStable(bool stable);
917
920 void setPCMaxCondSetSize(Size max_k);
921
925 void setPCUnshieldedColliderSorted(bool sorted);
926
928
929 // ##########################################################################
931 // ##########################################################################
933
936 void useFCIChi2Test();
937
940 void useFCIG2Test();
941
944 void setFCIAlpha(double alpha);
945
948 void setFCIMaxPathLength(Size max_len);
949
952 void setFCIExhaustiveSepSet(bool exhaustive);
953
956 bool fciExhaustiveSepSet() const;
957
959
960 // ##########################################################################
962 // ##########################################################################
964
966 void setMaxIndegree(Size max_indegree);
967
973 void setSliceOrder(const NodeProperty< NodeId >& slice_order);
974
981 void setSliceOrder(const std::vector< std::vector< std::string > >& slices);
982
988 void unsetSliceOrder();
989
1007 void setTotalOrder(const Sequence< NodeId >& order);
1008 void setTotalOrder(const std::vector< std::string >& order);
1010
1012 void unsetTotalOrder();
1013
1015
1017 void setForbiddenArcs(const ArcSet& set);
1018
1021 void addForbiddenArc(const Arc& arc);
1022 void addForbiddenArc(NodeId tail, NodeId head);
1023 void addForbiddenArc(std::string_view tail, std::string_view head);
1025
1028 void eraseForbiddenArc(const Arc& arc);
1029 void eraseForbiddenArc(NodeId tail, NodeId head);
1030 void eraseForbiddenArc(std::string_view tail, std::string_view head);
1032
1034 void setMandatoryArcs(const ArcSet& set);
1035
1038 void addMandatoryArc(const Arc& arc);
1039 void addMandatoryArc(NodeId tail, NodeId head);
1040 void addMandatoryArc(std::string_view tail, std::string_view head);
1042
1045 void eraseMandatoryArc(const Arc& arc);
1046 void eraseMandatoryArc(NodeId tail, NodeId head);
1047 void eraseMandatoryArc(std::string_view tail, std::string_view head);
1049
1052 void addNoParentNode(NodeId node);
1053 void addNoParentNode(std::string_view node);
1055
1058 void eraseNoParentNode(NodeId node);
1059 void eraseNoParentNode(std::string_view node);
1060
1063 void addNoChildrenNode(NodeId node);
1064 void addNoChildrenNode(std::string_view node);
1066
1069 void eraseNoChildrenNode(NodeId node);
1070 void eraseNoChildrenNode(std::string_view node);
1072
1077 void setPossibleEdges(const EdgeSet& set);
1078 void setPossibleSkeleton(const UndiGraph& skeleton);
1080
1085 void addPossibleEdge(const Edge& edge);
1086 void addPossibleEdge(NodeId tail, NodeId head);
1087 void addPossibleEdge(std::string_view tail, std::string_view head);
1089
1092 void erasePossibleEdge(const Edge& edge);
1093 void erasePossibleEdge(NodeId tail, NodeId head);
1094 void erasePossibleEdge(std::string_view tail, std::string_view head);
1096
1098
1099 // ##########################################################################
1101 // ##########################################################################
1103
1105
1109 void setNumberOfThreads(Size nb) override;
1110
1112
1113 protected:
1114 PAG learnPAG_();
1115 PDAG learnPDAG_();
1116
1118 void _setPriorWeight_(double weight);
1119
1121 bool inducedTypes_{false};
1122
1125
1127 Score* score_{nullptr};
1128
1131
1133 bool useEM_{false};
1134
1136 double noiseEM_{0.1};
1137
1140
1143
1145 Prior* prior_{nullptr};
1146
1148
1150 double priorWeight_{1.0f};
1151
1154
1157
1160
1163
1166
1169
1172
1175
1178
1181
1184
1187
1190
1193
1196
1199
1202
1206
1209
1211 enum class IndepTestType { Chi2, G2 };
1213
1216
1218 double alphaPc_{0.05};
1219 bool stablePc_{true};
1221 bool sortedUCPc_{false};
1222
1225
1228
1231
1233 double alphaFci_{0.05};
1236
1239
1242
1245
1248
1251
1254
1256 std::vector< std::pair< std::size_t, std::size_t > > ranges_;
1257
1260
1262 std::string priorDbname_;
1263
1266
1268 std::string filename_{"-"};
1269
1270 // size of the tabu list
1272
1273 // the current algorithm as an approximationScheme
1275
1277 static DatabaseTable readFile_(std::string_view filename,
1278 const std::vector< std::string >& missing_symbols);
1279
1281 static void isCSVFileName_(std::string_view filename);
1282
1284 virtual void createPrior_() = 0;
1285
1287 void createScore_();
1288
1291 bool take_into_account_score = true);
1292
1294 DAG learnDag_();
1295
1298
1301
1304
1307
1309 PriorType getPriorType_() const;
1310
1313
1314 public:
1315 // ##########################################################################
1318 // ##########################################################################
1319 // in order to not pollute the proper code of IBNLearner, we
1320 // directly
1321 // implement those
1322 // very simples methods here.
1324 void setCurrentApproximationScheme(const ApproximationScheme* approximationScheme);
1325
1326 void distributeProgress(const ApproximationScheme* approximationScheme,
1327 Size pourcent,
1328 double error,
1329 double time);
1330
1332 void distributeStop(const ApproximationScheme* approximationScheme, std::string_view message);
1333
1335
1340 void setEpsilon(double eps) override;
1341
1343 double epsilon() const override;
1344
1346 void disableEpsilon() override;
1347
1349 void enableEpsilon() override;
1350
1353 bool isEnabledEpsilon() const override;
1354
1356
1362 void setMinEpsilonRate(double rate) override;
1363
1365 double minEpsilonRate() const override;
1366
1368 void disableMinEpsilonRate() override;
1369
1371 void enableMinEpsilonRate() override;
1372
1375 bool isEnabledMinEpsilonRate() const override;
1376
1378
1384 void setMaxIter(Size max) override;
1385
1387 Size maxIter() const override;
1388
1390 void disableMaxIter() override;
1391
1393 void enableMaxIter() override;
1394
1397 bool isEnabledMaxIter() const override;
1398
1400
1405
1407 void setMaxTime(double timeout) override;
1408
1410 double maxTime() const override;
1411
1413 double currentTime() const override;
1414
1416 void disableMaxTime() override;
1417
1418 void enableMaxTime() override;
1419
1422 bool isEnabledMaxTime() const override;
1423
1425
1429 void setPeriodSize(Size p) override;
1430
1431 Size periodSize() const override;
1432
1434
1437 void setVerbosity(bool v) override;
1438
1439 bool verbosity() const override;
1440
1442
1445
1447
1449 Size nbrIterations() const override;
1450
1452 const std::vector< double >& history() const override;
1453
1455
1456
1459
1467 void EMsetEpsilon(double eps);
1468
1470
1475 double EMEpsilon() const;
1476
1478 void EMdisableEpsilon();
1479
1484 void EMenableEpsilon();
1485
1487 bool EMisEnabledEpsilon() const;
1488
1495 void EMsetMinEpsilonRate(double rate);
1496
1502 double EMMinEpsilonRate() const;
1503
1506
1512
1514 bool EMisEnabledMinEpsilonRate() const;
1515
1521 void EMsetMaxIter(Size max);
1522
1528 Size EMMaxIter() const;
1529
1531 void EMdisableMaxIter();
1532
1534 void EMenableMaxIter();
1535
1538 bool EMisEnabledMaxIter() const;
1539
1545 void EMsetMaxTime(double timeout);
1546
1552 double EMMaxTime() const;
1553
1555 double EMCurrentTime() const;
1556
1558 void EMdisableMaxTime();
1559
1560 void EMenableMaxTime();
1561
1563 bool EMisEnabledMaxTime() const;
1564
1569 void EMsetPeriodSize(Size p);
1570
1571 Size EMPeriodSize() const;
1572
1574 void EMsetVerbosity(bool v);
1575
1577 bool EMVerbosity() const;
1578
1581
1583 Size EMnbrIterations() const;
1584
1589 const std::vector< double >& EMHistory() const;
1590
1592 };
1593
1594 /* namespace learning */
1595} // namespace gum::learning
1596
1598#ifndef GUM_NO_INLINE
1600#endif /* GUM_NO_INLINE */
1601
1603
1604#endif /* GUM_LEARNING_GENERIC_BN_LEARNER_H */
A class that, given a structure and a parameter estimator returns a full Bayes net.
The class for initializing DatabaseTable and RawDatabaseTable instances from CSV files.
A DBRowGenerator class that returns the rows that are complete (fully observed) w....
A DBRowGenerator class that returns incomplete rows as EM would do.
A dirichlet priori: computes its N'_ijk from a database.
FCI (Fast Causal Inference) causal discovery algorithm.
A pack of learning algorithms that can easily be used.
The K2 algorithm.
PC (Peter-Clark) constraint-based structure learning algorithm.
The SimpleMiic algorithm.
Class representing a Bayesian network.
Definition BayesNet.h:99
Static math utilities for the chi2 distribution.
Definition chi2.h:73
Base class for dag.
Definition DAG.h:121
ApproximationSchemeSTATE
The different state of an approximation scheme.
Base class for mixed graphs.
Definition mixedGraph.h:146
Partial Ancestral Graph: undirected topology with endpoint marks.
Definition PAG.h:90
Base class for partially directed acyclic graphs.
Definition PDAG.h:130
ThreadNumberManager(Size nb_threads=0)
default constructor
A class that redirects gum_signal from algorithms to the listeners of BNLearn.
The class computing n times the corrected mutual information, as used in the MIIC algorithm.
KModeTypes
the description type for the complexity correction
A class that, given a structure and a parameter estimator returns a full Bayes net.
the class used to read a row in the database and to transform it into a set of DBRow instances that c...
The class representing a tabular database as used by learning tasks.
Fast Causal Inference — PAG learning via constraint-based methods.
Definition FCI.h:79
The greedy hill climbing learning algorithm (for directed graphs).
The greedy thick-thinning learning algorithm (for directed graphs).
a helper to easily read databases
Definition IBNLearner.h:140
const std::vector< std::string > & missingSymbols() const
returns the set of missing symbols taken into account
const DatabaseTable & databaseTable() const
returns the internal database table
Size _min_nb_rows_per_thread_
the minimal number of rows to parse (on average) by thread
Definition IBNLearner.h:291
std::size_t size() const
returns the number of records in the database
std::vector< std::size_t > _domain_sizes_
the domain sizes of the variables (useful to speed-up computations)
Definition IBNLearner.h:282
Database(std::string_view filename, const gum::BayesNet< GUM_SCALAR > &bn, const std::vector< std::string > &missing_symbols)
constructor with a BN providing the variables of interest
DatabaseTable _database_
the database itself
Definition IBNLearner.h:276
const std::string & nameFromId(NodeId id) const
returns the variable name corresponding to a given node id
double weight(const std::size_t i) const
returns the weight of the ith record
Bijection< NodeId, std::size_t > _nodeId2cols_
a bijection assigning to each variable name its NodeId
Definition IBNLearner.h:285
Database(std::string_view file, const std::vector< std::string > &missing_symbols, const bool induceTypes=false)
default constructor
const std::vector< std::string > & names() const
returns the names of the variables in the database
void setWeight(const std::size_t i, const double weight)
sets the weight of the ith record
const Bijection< NodeId, std::size_t > & nodeId2Columns() const
returns the mapping between node ids and their columns in the database
Database & operator=(const Database &from)
copy operator
DBRowGeneratorParser & parser()
returns the parser for the database
DBRowGeneratorParser * _parser_
the parser used for reading the database
Definition IBNLearner.h:279
NodeId idFromName(std::string_view var_name) const
returns the node id corresponding to a variable name
void setDatabaseWeight(const double new_weight)
assign a weight to all the rows of the database so that the sum of their weights is equal to new_weig...
Size _max_threads_number_
the max number of threads authorized
Definition IBNLearner.h:288
std::size_t nbRows() const
returns the number of records in the database
const std::vector< std::size_t > & domainSizes() const
returns the domain sizes of the variables
A pack of learning algorithms that can easily be used.
Definition IBNLearner.h:106
StructuralConstraintPossibleEdges constraintPossibleEdges_
the constraint on possible Edges
MixedGraph preparePC_()
prepares the initial graph and independence test for PC
StructuralConstraintNoParentNodes constraintNoParentNodes_
the constraint on no parent nodes
Size periodSize() const override
how many samples between 2 stopping isEnableds
BNLearnerPriorType priorType_
the a priorselected for the score and parameters
void EMenableEpsilon()
Enable the log-likelihood min diff stopping criterion in EM.
bool EMisEnabledEpsilon() const
return true if EM's stopping criterion is the log-likelihood min diff
void EMsetPeriodSize(Size p)
how many samples between 2 stoppings isEnabled
Size nbrIterations() const override
void setMinEpsilonRate(double rate) override
Given that we approximate f(t), stopping criterion on d/dt(|f(t+1)-f(t)|) If the criterion was disabl...
Size EMnbrIterations() const
returns the number of iterations performed by the last EM execution
void enableMaxTime() override
stopping criterion on timeout If the criterion was disabled it will be enabled
std::string priorDbname_
the filename for the Dirichlet a priori, if any
IndepTestType
independence test type for PC
double priorWeight_
the weight of the prior
void enableEpsilon() override
Enable stopping criterion on epsilon.
double maxTime() const override
returns the timeout (in seconds)
double noiseEM_
the noise factor (in (0,1)) used by EM for perturbing the CPT during init
std::vector< std::pair< std::size_t, std::size_t > > ranges_
the set of rows' ranges within the database in which learning is done
GreedyHillClimbing extendedGreedyHillClimbing_
the extended greedy hill climbing
ParamEstimatorType
an enumeration to select the type of parameter estimation we shall apply
Definition IBNLearner.h:113
AlgoType
an enumeration to select easily the learning algorithm to use
Definition IBNLearner.h:125
void distributeStop(const ApproximationScheme *approximationScheme, std::string_view message)
distribute signals
bool allowArcTriangleDeletions_
whether we allow or not arc deletions during learning
void EMdisableMinEpsilonRate()
Disable the log-likelihood evolution rate stopping criterion.
double EMMaxTime() const
@brief returns EM's timeout (in milliseconds)
virtual void createPrior_()=0
create the prior used for learning
ApproximationSchemeSTATE EMStateApproximationScheme() const
get the current state of EM
const std::vector< double > & history() const override
void setMaxIter(Size max) override
stopping criterion on number of iterationsIf the criterion was disabled it will be enabled
K2 algoK2_
the K2 algorithm
IndepTestType indepTestTypeFCI_
independence test type for FCI (reuses IndepTestType defined above)
AlgoType selectedAlgo_
the selected learning algorithm
void EMenableMaxIter()
Enable stopping criterion on max iterations.
void enableMinEpsilonRate() override
Enable stopping criterion on epsilon rate.
bool allowArcAdditions_
whether we allow or not arc additions during learning
void setEpsilon(double eps) override
Given that we approximate f(t), stopping criterion on |f(t+1)-f(t)| If the criterion was disabled it ...
void EMdisableEpsilon()
Disable the min log-likelihood diff stopping criterion for EM.
bool isEnabledMaxIter() const override
void EMsetMaxIter(Size max)
add a max iteration stopping criterion
Database scoreDatabase_
the database to be used by the scores and parameter estimators
void setMaxTime(double timeout) override
stopping criterion on timeout If the criterion was disabled it will be enabled
double epsilon() const override
Get the value of epsilon.
ScoreType
an enumeration enabling to select easily the score we wish to use
Definition IBNLearner.h:109
double EMEpsilon() const
Get the value of EM's min diff epsilon.
bool useEM_
a Boolean indicating whether we should use EM for parameter learning or not
DAG2BNLearner dag2BN_
the parametric EM
Prior * prior_
the prior used
void EMsetMinEpsilonRate(double rate)
sets the stopping criterion of EM as being the minimal log-likelihood's evolution rate
bool isEnabledMaxTime() const override
double EMMinEpsilonRate() const
Get the value of the minimal log-likelihood evolution rate of EM.
void distributeProgress(const ApproximationScheme *approximationScheme, Size pourcent, double error, double time)
{@ /// distribute signals
void setPeriodSize(Size p) override
how many samples between 2 stopping isEnableds
Size EMMaxIter() const
return the max number of iterations criterion
CorrectedMutualInformation * mutualInfo_
the selected correction for miic
void disableMinEpsilonRate() override
Disable stopping criterion on epsilon rate.
void EMdisableMaxIter()
Disable stopping criterion on max iterations.
BNLearnerPriorType
an enumeration to select the prior
Definition IBNLearner.h:116
bool isEnabledMinEpsilonRate() const override
const std::vector< double > & EMHistory() const
returns the history of the last EM execution
StructuralConstraintNoChildrenNodes constraintNoChildrenNodes_
the constraint on no children nodes
gum::learning::FCI algoFCI_
the FCI algorithm
void EMsetVerbosity(bool v)
sets or unsets EM's verbosity
Size maxIter() const override
ParamEstimatorType paramEstimatorType_
the type of the parameter estimator
ScoreType scoreType_
the score selected for learning
bool EMVerbosity() const
returns the EM's verbosity status
std::pair< double, double > G2(NodeId id1, NodeId id2, const std::vector< NodeId > &knowing={})
Return the <statistic,pvalue> pair for for G2 test in the database.
const ApproximationScheme * currentAlgorithm_
IndependenceTest * indepTestPC_
owned independence test object for PC (rebuilt before each learn call)
bool isEnabledEpsilon() const override
DAG learnDag_()
returns the DAG learnt
Database * priorDatabase_
the database used by the Dirichlet a priori
void createScore_()
create the score used for learning
PriorType getPriorType_() const
returns the type (as a string) of a given prior
double alphaFci_
FCI parameters.
double alphaPc_
PC parameters.
StructuralConstraintIndegree constraintIndegree_
the constraint for indegrees
bool allowArcDeletions_
whether we allow or not arc deletions during learning
bool verbosity() const override
verbosity
void disableMaxTime() override
Disable stopping criterion on timeout.
double currentTime() const override
get the current running time in second (double)
void disableEpsilon() override
Disable stopping criterion on epsilon.
std::string filename_
the filename database
void disableMaxIter() override
Disable stopping criterion on max iterations.
SimpleMiic algoSimpleMiic_
the MIIC algorithm
Score * score_
the score used
StructuralConstraintMandatoryArcs constraintMandatoryArcs_
the constraint on mandatory arcs
Miic algoMiic_
the Constraint MIIC algorithm
void EMenableMaxTime()
sets the stopping criterion of EM as being the minimal difference between two consecutive log-likelih...
void createCorrectedMutualInformation_()
create the Corrected Mutual Information instance for Miic
IndepTestType indepTestTypePC_
StructuralConstraintForbiddenArcs constraintForbiddenArcs_
the constraint on forbidden arcs
StructuralConstraintTotalOrder constraintTotalOrder_
the total order ing constraint
GreedyHillClimbing greedyHillClimbing_
the greedy hill climbing algorithm
void setVerbosity(bool v) override
verbosity
StructuralConstraintTabuList constraintTabuList_
the constraint for tabu lists
DAG initialDag_
an initial DAG given to learners
GreedyThickThinning greedyThickThinning_
the greedy thick-thinning algorithm
void EMsetMaxTime(double timeout)
add a stopping criterion on timeout
MixedGraph prepareFCI_()
prepares the initial graph and independence test for FCI
ApproximationSchemeSTATE stateApproximationScheme() const override
history
MixedGraph prepareSimpleMiic_()
prepares the initial graph for Simple Miic
gum::learning::PC algoPC_
the PC algorithm
void EMenableMinEpsilonRate()
Enable the log-likelihood evolution rate stopping criterion.
bool allowArcReversals_
whether we allow or not arc reversals during learning
Size EMPeriodSize() const
sets the stopping criterion of EM as being the minimal difference between two consecutive log-likelih...
double EMCurrentTime() const
get the current running time in second (double)
MixedGraph prepareMiic_()
prepares the initial graph for miic
IBNLearner(std::string_view filename, const std::vector< std::string > &missingSymbols, bool induceTypes=true)
read the database file for the score / parameter estimation and var names
LocalSearchWithTabuList localSearchWithTabuList_
the local search with tabu list algorithm
IndependenceTest * indepTestFCI_
owned independence test object for FCI (rebuilt before each learn call)
void setCurrentApproximationScheme(const ApproximationScheme *approximationScheme)
{@ /// distribute signals
ParamEstimator * createParamEstimator_(const DBRowGeneratorParser &parser, bool take_into_account_score=true)
create the parameter estimator used for learning
StructuralConstraintSliceOrder constraintSliceOrder_
the constraint for 2TBNs
bool EMisEnabledMinEpsilonRate() const
static constexpr double default_EM_noise
the default noise amount added to CPTs during EM's initialization (see method useEM())
Definition IBNLearner.h:137
void EMdisableMaxTime()
Disable EM's timeout stopping criterion.
void EMsetEpsilon(double eps)
sets the stopping criterion of EM as being the minimal difference between two consecutive log-likelih...
void enableMaxIter() override
Enable stopping criterion on max iterations.
bool inducedTypes_
the policy for typing variables
CorrectedMutualInformation::KModeTypes kmodeMiic_
the penalty used in MIIC
double minEpsilonRate() const override
Get the value of the minimal epsilon rate.
The base class for all the independence tests used for learning.
The K2 algorithm.
Definition K2.h:63
The local search with tabu list learning algorithm (for directed graphs).
The MIIC learning algorithm.
Definition Miic.h:100
the no a priorclass: corresponds to 0 weight-sample
Definition noPrior.h:65
PC (Peter-Clark) constraint-based structure learning algorithm.
Definition PC.h:74
The base class for estimating parameters of CPTs.
the base class for all a priori
Definition prior.h:84
The base class for all the scores used for learning (BIC, BDeu, etc).
Definition score.h:68
The miic learning algorithm.
Definition SimpleMiic.h:83
the structural constraint for forbidding the creation of some arcs during structure learning
the class for structural constraints limiting the number of parents of nodes in a directed graph
the structural constraint indicating that some arcs shall never be removed or reversed
the structural constraint for forbidding children for some nodes
the structural constraint for forbidding parents for some nodes
the structural constraint for forbidding the creation of some arcs except those defined in the class ...
the structural constraint imposing a partial order over nodes
The class imposing a N-sized tabu list as a structural constraints for learning algorithms.
the structural constraint imposing a total order over some nodes
Class building the essential Graph from a DAGmodel.
The basic class for computing the set of digraph changes allowed by the user to be executed by the le...
The basic class for computing the set of digraph changes allowed by the user to be executed by the le...
The mecanism to compute the next available graph changes for directed structure learning search algor...
The greedy thick-thinning learning algorithm (for directed graphs).
void setNumberOfThreads(unsigned int number)
Set the max number of threads to be used when entering the next parallel region.
Definition threads.cpp:55
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Definition types.h:74
Size NodeId
Type for node ids.
the class for computing Chi2 scores
the class for computing G2 scores
The local search learning with tabu list algorithm (for directed graphs).
include the inlined functions if necessary
Definition CSVParser.h:55
unsigned int getNumberOfThreads()
returns the max number of threads used by default when entering the next parallel region
the class for estimating parameters of CPTs using Maximum Likelihood
the class for computing AIC scores
the class for computing Bayesian Dirichlet (BD) log2 scores
the class for computing BDeu scores
the class for computing K2 scores (actually their log2 value)
the class for computing fNML scores
the base class for structural constraints imposed by DAGs
the structural constraint for forbidding the creation of some arcs during structure learning
the class for structural constraints limiting the number of parents of nodes in a directed graph
the structural constraint indicating that some arcs shall never be removed or reversed
the structural constraint for forbidding children for some nodes during structure learning
the structural constraint for forbidding parents for some nodes during structure learning
the structural constraint for forbidding the creation of some arcs during structure learning
the structural constraint imposing a partial order over nodes
the class imposing a N-sized tabu list as a structural constraints for learning algorithms
the structural constraint imposing a total ordering over some nodes