56#ifndef DOXYGEN_SHOULD_SKIP_THIS
65 template < GUM_Numeric GUM_SCALAR >
67 const std::vector< std::string >& missingSymbols,
68 const bool induceTypes) :
69 IBNLearner(filename, missingSymbols, induceTypes) {
70 GUM_CONSTRUCTOR(BNLearner);
73 template < GUM_Numeric GUM_SCALAR >
74 BNLearner< GUM_SCALAR >::BNLearner(
const DatabaseTable& db) : IBNLearner(db) {
75 GUM_CONSTRUCTOR(BNLearner);
78 template < GUM_Numeric GUM_SCALAR >
79 BNLearner< GUM_SCALAR >::BNLearner(std::string_view filename,
81 const std::vector< std::string >& missing_symbols) :
82 IBNLearner(filename, bn, missing_symbols) {
83 GUM_CONSTRUCTOR(BNLearner);
87 template < GUM_Numeric GUM_SCALAR >
88 BNLearner< GUM_SCALAR >::BNLearner(
const BNLearner< GUM_SCALAR >& src) : IBNLearner(src) {
89 GUM_CONSTRUCTOR(BNLearner);
93 template < GUM_Numeric GUM_SCALAR >
94 BNLearner< GUM_SCALAR >::BNLearner(BNLearner< GUM_SCALAR >&& src) : IBNLearner(src) {
95 GUM_CONSTRUCTOR(BNLearner);
99 template < GUM_Numeric GUM_SCALAR >
100 BNLearner< GUM_SCALAR >::~BNLearner() {
112 template < GUM_Numeric GUM_SCALAR >
120 template < GUM_Numeric GUM_SCALAR >
128 template < GUM_Numeric GUM_SCALAR >
131 auto notification = checkScorePriorCompatibility();
132 if (notification !=
"") { std::cout <<
"[aGrUM notification] " << notification << std::endl; }
136 std::unique_ptr< ParamEstimator > param_estimator(
137 createParamEstimator_(scoreDatabase_.parser(),
true));
139 return dag2BN_.createBN< GUM_SCALAR >(*(param_estimator.get()), learnDag_());
143 template < GUM_Numeric GUM_SCALAR >
146 if (dag.size() == 0)
return;
149 std::vector< NodeId > ids;
150 ids.reserve(dag.sizeNodes());
151 for (
const auto node: dag)
153 std::sort(ids.begin(), ids.end());
155 if (ids.back() >= scoreDatabase_.names().size()) {
156 std::string str =
"Learning parameters corresponding to the dag is impossible "
157 "because the database does not contain the following nodeID";
158 std::vector< NodeId > bad_ids;
159 for (
const auto node: ids) {
160 if (node >= scoreDatabase_.names().size()) bad_ids.push_back(node);
162 if (bad_ids.size() > 1) str +=
's';
165 for (
const auto node: bad_ids) {
166 if (deja) str +=
", ";
168 str += std::to_string(node);
170 GUM_ERROR(MissingVariableInDatabase, str)
175 template < GUM_Numeric GUM_SCALAR >
177 bool takeIntoAccountScore) {
179 if (dag.size() == 0)
return BayesNet< GUM_SCALAR >();
182 _checkDAGCompatibility_(dag);
188 if (scoreDatabase_.databaseTable().hasMissingValues()
189 || ((priorDatabase_ !=
nullptr)
190 && (priorType_ == BNLearnerPriorType::DIRICHLET_FROM_DATABASE)
191 && priorDatabase_->databaseTable().hasMissingValues())) {
193 "In general, the BNLearner is unable to cope with "
194 <<
"missing values in databases. To learn parameters in "
195 <<
"such situations, you should first use method " <<
"useEM()");
200 std::unique_ptr< ParamEstimator > param_estimator(
201 createParamEstimator_(parser, takeIntoAccountScore));
203 return dag2BN_.createBN< GUM_SCALAR >(*(param_estimator.get()), dag);
211# if defined(__GNUC__) && !defined(__clang__)
212# pragma GCC push_options
213# pragma GCC optimize("no-tree-vrp")
216 template < GUM_Numeric GUM_SCALAR >
217 std::pair< std::shared_ptr< ParamEstimator >, std::shared_ptr< ParamEstimator > >
219 bool takeIntoAccountScore) {
221 _checkDAGCompatibility_(dag);
231 const auto& database = scoreDatabase_.databaseTable();
232 const std::size_t nb_vars = database.nbVariables();
233 const std::vector< gum::learning::DBTranslatedValueType > col_types(
242 std::shared_ptr< ParamEstimator > param_estimator_bootstrap(
243 createParamEstimator_(parser_bootstrap, takeIntoAccountScore));
246 BayesNet< GUM_SCALAR > dummy_bn;
252 std::shared_ptr< ParamEstimator > param_estimator_EM(
253 createParamEstimator_(parser_EM, takeIntoAccountScore));
255 return {param_estimator_bootstrap, param_estimator_EM};
258# if defined(__GNUC__) && !defined(__clang__)
259# pragma GCC pop_options
263 template < GUM_Numeric GUM_SCALAR >
264 BayesNet< GUM_SCALAR >
266 bool takeIntoAccountScore) {
268 if (dag.size() == 0)
return BayesNet< GUM_SCALAR >();
271 auto estimators = _initializeEMParameterLearning_(dag, takeIntoAccountScore);
274 return dag2BN_.createBNwithEM< GUM_SCALAR >(*(estimators.first.get()),
275 *(estimators.second.get()),
280 template < GUM_Numeric GUM_SCALAR >
281 BayesNet< GUM_SCALAR >
283 bool takeIntoAccountScore) {
285 if (bn.internalDag().size() == 0)
return BayesNet< GUM_SCALAR >();
288 auto estimators = _initializeEMParameterLearning_(bn.internalDag(), takeIntoAccountScore);
290 return dag2BN_.createBNwithEM< GUM_SCALAR >(*(estimators.first.get()),
291 *(estimators.second.get()),
296 template < GUM_Numeric GUM_SCALAR >
298 bool takeIntoAccountScore) {
299 if (!scoreDatabase_.databaseTable().hasMissingValues() || !useEM_) {
301 return _learnParameters_(dag, takeIntoAccountScore);
304 return _learnParametersWithEM_(dag, takeIntoAccountScore);
309 template < GUM_Numeric GUM_SCALAR >
310 BayesNet< GUM_SCALAR >
312 bool takeIntoAccountScore) {
313 if (!scoreDatabase_.databaseTable().hasMissingValues() || !useEM_) {
315 const auto& db = scoreDatabase_.databaseTable();
316 for (
const auto n: bn.nodes()) {
317 dag.addNodeWithId(db.columnFromVariableName(bn.variable(n).name()));
319 for (
const auto& arc: bn.arcs()) {
320 dag.addArc(db.columnFromVariableName(bn.variable(arc.tail()).name()),
321 db.columnFromVariableName(bn.variable(arc.head()).name()));
325 return _learnParameters_(dag, takeIntoAccountScore);
327 return _learnParametersWithEM_(bn, takeIntoAccountScore);
332 template < GUM_Numeric GUM_SCALAR >
334 return learnParameters(initialDag_, take_into_account_score);
337 template < GUM_Numeric GUM_SCALAR >
340 const BayesNet< GUM_SCALAR >& src) {
341 std::ifstream in(std::string(filename), std::ifstream::in);
343 if ((in.rdstate() & std::ifstream::failbit) != 0) {
344 GUM_ERROR(gum::IOError,
"File " << filename <<
" not found")
347 CSVParser parser(in, std::string(filename));
349 auto names = parser.current();
353 for (
gum::Idx col = 0; col < names.size(); col++) {
354 if (src.exists(names[col])) {
356 modals.insert(col, gum::Sequence< std::string >());
358 for (
gum::Size i = 0; i < src.variable(graphId).domainSize(); ++i)
359 modals[col].insert(src.variable(graphId).label(i));
367 template < GUM_Numeric GUM_SCALAR >
369 const auto st = state();
372 for (
const auto& tuple: st)
373 if (std::get< 0 >(tuple).length() > maxkey) maxkey = std::get< 0 >(tuple).length();
376 for (
const auto& tuple: st) {
377 s += std::format(
"{:<{}} : {}", std::get< 0 >(tuple), maxkey, std::get< 1 >(tuple));
378 if (std::get< 2 >(tuple) !=
"") s += std::format(
" ({})", std::get< 2 >(tuple));
384 template < GUM_Numeric GUM_SCALAR >
385 std::vector< std::tuple< std::string, std::string, std::string > >
387 std::vector< std::tuple< std::string, std::string, std::string > > vals;
391 const auto& db = database();
393 vals.emplace_back(
"Filename", filename_,
"");
394 vals.emplace_back(
"Size",
395 "(" + std::to_string(nbRows()) +
"," + std::to_string(nbCols()) +
")",
398 std::string vars =
"";
399 for (
NodeId i = 0; i < db.nbVariables(); i++) {
400 if (i > 0) vars +=
", ";
401 vars += nameFromId(i) +
"[" + std::to_string(db.domainSize(i)) +
"]";
403 vals.emplace_back(
"Variables", vars,
"");
404 vals.emplace_back(
"Induced types", inducedTypes_ ?
"True" :
"False",
"");
405 vals.emplace_back(
"Missing values", hasMissingValues() ?
"True" :
"False",
"");
408 switch (selectedAlgo_) {
409 case AlgoType::GREEDY_HILL_CLIMBING :
410 vals.emplace_back(key,
"Greedy Hill Climbing",
"");
412 case AlgoType::EXTENDED_GREEDY_HILL_CLIMBING :
413 vals.emplace_back(key,
"Extended Greedy Hill Climbing",
"");
415 case AlgoType::K2 : {
416 vals.emplace_back(key,
"K2",
"");
417 const auto& k2order = algoK2_.order();
419 for (
NodeId i = 0; i < k2order.size(); i++) {
420 if (i > 0) vars +=
", ";
421 vars += nameFromId(k2order.atPos(i));
423 vals.emplace_back(
"K2 order", vars,
"");
425 case AlgoType::LOCAL_SEARCH_WITH_TABU_LIST :
426 vals.emplace_back(key,
"Local Search with Tabu List",
"");
427 vals.emplace_back(
"Tabu list size", std::to_string(nbDecreasingChanges_),
"");
429 case AlgoType::MIIC : vals.emplace_back(key,
"MIIC",
"");
break;
430 case AlgoType::PC : vals.emplace_back(key,
"PC",
"");
break;
431 case AlgoType::FCI : vals.emplace_back(key,
"FCI",
"");
break;
432 case AlgoType::GREEDY_THICK_THINNING :
433 vals.emplace_back(key,
"Greedy Thick Thinning",
"");
435 default : vals.emplace_back(key,
"(unknown)",
"?");
break;
440 if (isScoreBased()) {
441 switch (scoreType_) {
442 case ScoreType::AIC : vals.emplace_back(key,
"AIC",
"");
break;
443 case ScoreType::BIC : vals.emplace_back(key,
"BIC",
"");
break;
444 case ScoreType::BD : vals.emplace_back(key,
"BD",
"");
break;
445 case ScoreType::BDeu : vals.emplace_back(key,
"BDeu",
"");
break;
446 case ScoreType::fNML : vals.emplace_back(key,
"fNML",
"");
break;
447 case ScoreType::K2 : vals.emplace_back(key,
"K2",
"");
break;
448 case ScoreType::LOG2LIKELIHOOD : vals.emplace_back(key,
"Log2Likelihood",
"");
break;
449 case ScoreType::MDL : vals.emplace_back(key,
"MDL",
"");
break;
450 default : vals.emplace_back(key,
"(unknown)",
"?");
break;
454 if (isConstraintBased()) {
456 switch (kmodeMiic_) {
458 vals.emplace_back(key,
"MDL",
"");
461 vals.emplace_back(key,
"NML",
"");
464 vals.emplace_back(key,
"No correction",
"");
466 default : vals.emplace_back(key,
"(unknown)",
"?");
break;
471 comment = checkScorePriorCompatibility();
472 switch (priorType_) {
473 case BNLearnerPriorType::NO_prior : vals.emplace_back(key,
"-", comment);
break;
474 case BNLearnerPriorType::DIRICHLET_FROM_DATABASE :
475 vals.emplace_back(key,
"Dirichlet", comment);
476 vals.emplace_back(
"Dirichlet from database", priorDbname_,
"");
478 case BNLearnerPriorType::DIRICHLET_FROM_BAYESNET :
479 vals.emplace_back(key,
"Dirichlet", comment);
480 vals.emplace_back(
"Dirichlet from Bayesian network : ", _prior_bn_.toString(),
"");
482 case BNLearnerPriorType::BDEU : vals.emplace_back(key,
"BDEU", comment);
break;
483 case BNLearnerPriorType::SMOOTHING : vals.emplace_back(key,
"Smoothing", comment);
break;
484 default : vals.emplace_back(key,
"(unknown)",
"?");
break;
487 if (priorType_ != BNLearnerPriorType::NO_prior)
488 vals.emplace_back(
"Prior weight", std::to_string(priorWeight_),
"");
490 if (databaseWeight() !=
double(nbRows())) {
491 vals.emplace_back(
"Database weight", std::to_string(databaseWeight()),
"");
496 if (!hasMissingValues()) comment =
"But no missing values in this database";
497 vals.emplace_back(
"use EM",
"True",
"");
500 if (dag2BN_.isEnabledMinEpsilonRate()) {
501 s += std::format(
"MinRate: {}", dag2BN_.minEpsilonRate());
504 if (dag2BN_.isEnabledEpsilon()) {
505 if (!first) s +=
", ";
507 s += std::format(
"MinDiff: {}", dag2BN_.epsilon());
509 if (dag2BN_.isEnabledMaxIter()) {
510 if (!first) s +=
", ";
512 s += std::format(
"MaxIter: {}", dag2BN_.maxIter());
514 if (dag2BN_.isEnabledMaxTime()) {
515 if (!first) s +=
", ";
517 s += std::format(
"MaxTime: {}", dag2BN_.maxTime());
520 vals.emplace_back(
"EM stopping criteria", s, comment);
525 if (constraintIndegree_.maxIndegree() < std::numeric_limits< Size >::max()) {
526 vals.emplace_back(
"Constraint Max InDegree",
527 std::to_string(constraintIndegree_.maxIndegree()),
530 if (!constraintForbiddenArcs_.arcs().empty()) {
533 for (
const auto& arc: constraintForbiddenArcs_.arcs()) {
534 if (nofirst) res +=
", ";
536 res += nameFromId(arc.tail()) +
"->" + nameFromId(arc.head());
539 vals.emplace_back(
"Constraint Forbidden Arcs", res,
"");
541 if (!constraintMandatoryArcs_.arcs().empty()) {
544 for (
const auto& arc: constraintMandatoryArcs_.arcs()) {
545 if (nofirst) res +=
", ";
547 res += nameFromId(arc.tail()) +
"->" + nameFromId(arc.head());
550 vals.emplace_back(
"Constraint Mandatory Arcs", res,
"");
552 if (!constraintPossibleEdges_.edges().empty()) {
555 for (
const auto& edge: constraintPossibleEdges_.edges()) {
556 if (nofirst) res +=
", ";
558 res += nameFromId(edge.first()) +
"--" + nameFromId(edge.second());
561 vals.emplace_back(
"Constraint Possible Edges", res,
"");
563 if (!constraintSliceOrder_.sliceOrder().empty()) {
566 const auto& order = constraintSliceOrder_.sliceOrder();
567 for (
const auto& p: order) {
568 if (nofirst) res +=
", ";
570 res += nameFromId(p.first) +
":" + std::to_string(p.second);
573 vals.emplace_back(
"Constraint Slice Order", res,
"");
575 if (!constraintNoParentNodes_.nodes().empty()) {
578 for (
const auto& node: constraintNoParentNodes_.nodes()) {
579 if (nofirst) res +=
", ";
581 res += nameFromId(node);
584 vals.emplace_back(
"Constraint No Parent Nodes", res,
"");
586 if (!constraintNoChildrenNodes_.nodes().empty()) {
589 for (
const auto& node: constraintNoChildrenNodes_.nodes()) {
590 if (nofirst) res +=
", ";
592 res += nameFromId(node);
595 vals.emplace_back(
"Constraint No Children Nodes", res,
"");
597 if (initialDag_.size() != 0) {
598 vals.emplace_back(
"Initial DAG",
"True", initialDag_.toDot());
604 template < GUM_Numeric GUM_SCALAR >
606 switch (learner.selectedAlgo_) {
607 case AlgoType::EXTENDED_GREEDY_HILL_CLIMBING : useExtendedGreedyHillClimbing();
break;
608 case AlgoType::GREEDY_HILL_CLIMBING : useGreedyHillClimbing();
break;
609 case AlgoType::GREEDY_THICK_THINNING :
610 useGreedyThickThinning();
611 setGreedyThickThinningReversals(learner.greedyThickThinningReversals());
613 case AlgoType::K2 : useK2(learner.algoK2_.order());
break;
614 case AlgoType::LOCAL_SEARCH_WITH_TABU_LIST :
615 useLocalSearchWithTabuList(learner.nbDecreasingChanges_);
617 case AlgoType::MIIC : useMIIC();
break;
618 case AlgoType::PC : usePC();
break;
619 case AlgoType::FCI : useFCI();
break;
622 switch (learner.scoreType_) {
623 case ScoreType::K2 : useScoreK2();
break;
624 case ScoreType::AIC : useScoreAIC();
break;
625 case ScoreType::BIC : useScoreBIC();
break;
626 case ScoreType::BD : useScoreBD();
break;
627 case ScoreType::BDeu : useScoreBDeu();
break;
628 case ScoreType::fNML : useScorefNML();
break;
629 case ScoreType::MDL : useScoreMDL();
break;
630 case ScoreType::LOG2LIKELIHOOD : useScoreLog2Likelihood();
break;
633 switch (learner.kmodeMiic_) {
639 switch (learner.priorType_) {
640 case BNLearnerPriorType::NO_prior : useNoPrior();
break;
641 case BNLearnerPriorType::DIRICHLET_FROM_DATABASE :
642 useDirichletPrior(learner.priorDbname_, learner.priorWeight_);
644 case BNLearnerPriorType::DIRICHLET_FROM_BAYESNET :
645 useDirichletPrior(learner._prior_bn_);
647 case BNLearnerPriorType::BDEU : useBDeuPrior(learner.priorWeight_);
break;
648 case BNLearnerPriorType::SMOOTHING : useSmoothingPrior(learner.priorWeight_);
break;
651 useEM_ = learner.useEM_;
652 noiseEM_ = learner.noiseEM_;
653 dag2BN_ = learner.dag2BN_;
655 setMaxIndegree(learner.constraintIndegree_.maxIndegree());
656 for (
const auto src: learner.constraintNoParentNodes_.nodes()) {
658 const auto dst = idFromName(learner.nameFromId(src));
659 addNoParentNode(dst);
660 }
catch (
const MissingVariableInDatabase&) {
664 for (
const auto src: learner.constraintNoChildrenNodes_.nodes()) {
666 const auto dst = idFromName(learner.nameFromId(src));
667 addNoChildrenNode(dst);
668 }
catch (
const MissingVariableInDatabase&) {
672 for (
const auto& arc: learner.constraintForbiddenArcs_.arcs()) {
674 const auto src = idFromName(learner.nameFromId(arc.tail()));
675 const auto dst = idFromName(learner.nameFromId(arc.head()));
676 addForbiddenArc(src, dst);
677 }
catch (
const MissingVariableInDatabase&) {
681 for (
const auto& arc: learner.constraintMandatoryArcs_.arcs()) {
683 const auto src = idFromName(learner.nameFromId(arc.tail()));
684 const auto dst = idFromName(learner.nameFromId(arc.head()));
685 addMandatoryArc(src, dst);
686 }
catch (
const MissingVariableInDatabase&) {
690 for (
const auto& edge: learner.constraintPossibleEdges_.edges()) {
692 const auto src = idFromName(learner.nameFromId(edge.first()));
693 const auto dst = idFromName(learner.nameFromId(edge.second()));
694 addPossibleEdge(src, dst);
695 }
catch (
const MissingVariableInDatabase&) {
699 if (!learner.constraintSliceOrder_.sliceOrder().empty()) {
701 for (
const auto& p: learner.constraintSliceOrder_.sliceOrder()) {
703 slice_order.insert(idFromName(learner.nameFromId(p.first)), p.second);
704 }
catch (
const MissingVariableInDatabase&) {
708 setSliceOrder(slice_order);
710 if (!learner.constraintTotalOrder_.totalOrder().empty()) {
711 setTotalOrder(learner.constraintTotalOrder_.totalOrder());
715 template < GUM_Numeric GUM_SCALAR >
718 Prior* old_prior = prior_;
721 switch (priorType_) {
722 case BNLearnerPriorType::NO_prior :
723 prior_ =
new NoPrior(scoreDatabase_.databaseTable(), scoreDatabase_.nodeId2Columns());
726 case BNLearnerPriorType::SMOOTHING :
728 =
new SmoothingPrior(scoreDatabase_.databaseTable(), scoreDatabase_.nodeId2Columns());
731 case BNLearnerPriorType::DIRICHLET_FROM_DATABASE :
732 if (priorDatabase_ !=
nullptr) {
733 delete priorDatabase_;
734 priorDatabase_ =
nullptr;
738 =
new Database(priorDbname_, scoreDatabase_, scoreDatabase_.missingSymbols());
741 priorDatabase_->parser(),
742 priorDatabase_->nodeId2Columns());
745 case BNLearnerPriorType::DIRICHLET_FROM_BAYESNET :
750 case BNLearnerPriorType::BDEU :
751 prior_ =
new BDeuPrior(scoreDatabase_.databaseTable(), scoreDatabase_.nodeId2Columns());
754 default :
GUM_ERROR(OperationNotAllowed,
"The BNLearner does not support yet this prior")
761 if (old_prior !=
nullptr)
delete old_prior;
764 template < GUM_Numeric GUM_SCALAR >
766 output << learner.toString();
774 template < GUM_Numeric GUM_SCALAR >
780 template < GUM_Numeric GUM_SCALAR >
782 const double noise) {
787 template < GUM_Numeric GUM_SCALAR >
789 const double noise) {
794 template < GUM_Numeric GUM_SCALAR >
796 const double noise) {
801 template < GUM_Numeric GUM_SCALAR >
807 template < GUM_Numeric GUM_SCALAR >
813 template < GUM_Numeric GUM_SCALAR >
819 template < GUM_Numeric GUM_SCALAR >
825 template < GUM_Numeric GUM_SCALAR >
831 template < GUM_Numeric GUM_SCALAR >
837 template < GUM_Numeric GUM_SCALAR >
843 template < GUM_Numeric GUM_SCALAR >
849 template < GUM_Numeric GUM_SCALAR >
855 template < GUM_Numeric GUM_SCALAR >
861 template < GUM_Numeric GUM_SCALAR >
867 template < GUM_Numeric GUM_SCALAR >
873 template < GUM_Numeric GUM_SCALAR >
879 template < GUM_Numeric GUM_SCALAR >
885 template < GUM_Numeric GUM_SCALAR >
891 template < GUM_Numeric GUM_SCALAR >
897 template < GUM_Numeric GUM_SCALAR >
903 template < GUM_Numeric GUM_SCALAR >
909 template < GUM_Numeric GUM_SCALAR >
915 template < GUM_Numeric GUM_SCALAR >
921 template < GUM_Numeric GUM_SCALAR >
927 template < GUM_Numeric GUM_SCALAR >
933 template < GUM_Numeric GUM_SCALAR >
939 template < GUM_Numeric GUM_SCALAR >
945 template < GUM_Numeric GUM_SCALAR >
952 template < GUM_Numeric GUM_SCALAR >
957 priorType_ = BNLearnerPriorType::DIRICHLET_FROM_BAYESNET;
958 _setPriorWeight_(weight);
962 template < GUM_Numeric GUM_SCALAR >
968 template < GUM_Numeric GUM_SCALAR >
974 template < GUM_Numeric GUM_SCALAR >
980 template < GUM_Numeric GUM_SCALAR >
986 template < GUM_Numeric GUM_SCALAR >
991 template < GUM_Numeric GUM_SCALAR >
998 template < GUM_Numeric GUM_SCALAR >
1004 template < GUM_Numeric GUM_SCALAR >
1010 template < GUM_Numeric GUM_SCALAR >
1016 template < GUM_Numeric GUM_SCALAR >
1022 template < GUM_Numeric GUM_SCALAR >
1028 template < GUM_Numeric GUM_SCALAR >
1034 template < GUM_Numeric GUM_SCALAR >
1040 template < GUM_Numeric GUM_SCALAR >
1046 template < GUM_Numeric GUM_SCALAR >
1052 template < GUM_Numeric GUM_SCALAR >
1058 template < GUM_Numeric GUM_SCALAR >
1063 template < GUM_Numeric GUM_SCALAR >
1069 template < GUM_Numeric GUM_SCALAR >
1075 template < GUM_Numeric GUM_SCALAR >
1081 template < GUM_Numeric GUM_SCALAR >
1087 template < GUM_Numeric GUM_SCALAR >
1093 template < GUM_Numeric GUM_SCALAR >
1099 template < GUM_Numeric GUM_SCALAR >
1105 template < GUM_Numeric GUM_SCALAR >
1111 template < GUM_Numeric GUM_SCALAR >
1117 template < GUM_Numeric GUM_SCALAR >
1123 template < GUM_Numeric GUM_SCALAR >
1130 template < GUM_Numeric GUM_SCALAR >
1132 const std::vector< std::vector< std::string > >& slices) {
1137 template < GUM_Numeric GUM_SCALAR >
1144 template < GUM_Numeric GUM_SCALAR >
1151 template < GUM_Numeric GUM_SCALAR >
1157 template < GUM_Numeric GUM_SCALAR >
1163 template < GUM_Numeric GUM_SCALAR >
1169 template < GUM_Numeric GUM_SCALAR >
1171 std::string_view head) {
1176 template < GUM_Numeric GUM_SCALAR >
1182 template < GUM_Numeric GUM_SCALAR >
1188 template < GUM_Numeric GUM_SCALAR >
1190 std::string_view head) {
1195 template < GUM_Numeric GUM_SCALAR >
1201 template < GUM_Numeric GUM_SCALAR >
1207 template < GUM_Numeric GUM_SCALAR >
1209 std::string_view head) {
1214 template < GUM_Numeric GUM_SCALAR >
1220 template < GUM_Numeric GUM_SCALAR >
1226 template < GUM_Numeric GUM_SCALAR >
1228 std::string_view head) {
1233 template < GUM_Numeric GUM_SCALAR >
1239 template < GUM_Numeric GUM_SCALAR >
1245 template < GUM_Numeric GUM_SCALAR >
1247 std::string_view head) {
1252 template < GUM_Numeric GUM_SCALAR >
1258 template < GUM_Numeric GUM_SCALAR >
1264 template < GUM_Numeric GUM_SCALAR >
1266 std::string_view head) {
1271 template < GUM_Numeric GUM_SCALAR >
1277 template < GUM_Numeric GUM_SCALAR >
1283 template < GUM_Numeric GUM_SCALAR >
1290 template < GUM_Numeric GUM_SCALAR >
1296 template < GUM_Numeric GUM_SCALAR >
1302 template < GUM_Numeric GUM_SCALAR >
1308 template < GUM_Numeric GUM_SCALAR >
1314 template < GUM_Numeric GUM_SCALAR >
1320 template < GUM_Numeric GUM_SCALAR >
1326 template < GUM_Numeric GUM_SCALAR >
1332 template < GUM_Numeric GUM_SCALAR >
1338 template < GUM_Numeric GUM_SCALAR >
1344 template < GUM_Numeric GUM_SCALAR >
1350 template < GUM_Numeric GUM_SCALAR >
1356 template < GUM_Numeric GUM_SCALAR >
1362 template < GUM_Numeric GUM_SCALAR >
1367 template < GUM_Numeric GUM_SCALAR >
A listener that allows BNLearner to be used as a proxy for its inner algorithms.
A basic pack of learning algorithms that can easily be used.
Class representing a Bayesian network.
BDeuPrior(const DatabaseTable &database, const Bijection< NodeId, std::size_t > &nodeId2columns=Bijection< NodeId, std::size_t >())
default constructor
A pack of learning algorithms that can easily be used.
BNLearner< GUM_SCALAR > & useNoPrior()
BNLearner< GUM_SCALAR > & EMdisableMaxIter()
Disable stopping criterion on max iterations.
BNLearner< GUM_SCALAR > & useFCIChi2Test()
BNLearner< GUM_SCALAR > & EMenableMinEpsilonRate()
Enable the log-likelihood evolution rate stopping criterion.
BNLearner< GUM_SCALAR > & useGreedyThickThinning()
BNLearner< GUM_SCALAR > & addForbiddenArc(const Arc &arc)
BNLearner< GUM_SCALAR > & useScoreAIC()
BNLearner< GUM_SCALAR > & eraseMandatoryArc(const Arc &arc)
BNLearner< GUM_SCALAR > & setTotalOrder(const std::vector< std::string > &order)
BNLearner< GUM_SCALAR > & EMdisableEpsilon()
Disable the min log-likelihood diff stopping criterion.
BNLearner< GUM_SCALAR > & addPossibleEdge(const Edge &edge)
BNLearner< GUM_SCALAR > & useScoreBD()
BNLearner< GUM_SCALAR > & allowArcReversals(bool allow)
std::vector< std::tuple< std::string, std::string, std::string > > state() const
BNLearner< GUM_SCALAR > & setPCStable(bool stable)
BNLearner< GUM_SCALAR > & addMandatoryArc(const Arc &arc)
BNLearner< GUM_SCALAR > & useBDeuPrior(double weight=1.0)
BNLearner< GUM_SCALAR > & useScoreBIC()
BNLearner< GUM_SCALAR > & usePC()
BNLearner< GUM_SCALAR > & useNoCorrection()
std::pair< std::shared_ptr< ParamEstimator >, std::shared_ptr< ParamEstimator > > _initializeEMParameterLearning_(const DAG &dag, bool takeIntoAccountScore)
initializes EM and returns a pair containing, first, a bootstrap estimator and, second,...
BNLearner< GUM_SCALAR > & useG2Test()
BNLearner< GUM_SCALAR > & setFCIMaxPathLength(Size max_len)
BNLearner< GUM_SCALAR > & useSmoothingPrior(double weight=1)
BNLearner< GUM_SCALAR > & EMdisableMinEpsilonRate()
Disable the log-likelihood evolution rate stopping criterion.
BNLearner(std::string_view filename, const std::vector< std::string > &missingSymbols={"?"}, const bool induceTypes=true)
default constructor
BNLearner< GUM_SCALAR > & erasePossibleEdge(const Edge &edge)
BNLearner< GUM_SCALAR > & setGreedyThickThinningReversals(bool allow)
BNLearner< GUM_SCALAR > & setPossibleSkeleton(const UndiGraph &skeleton)
BayesNet< GUM_SCALAR > _learnParameters_(const DAG &dag, bool takeIntoAccountScore)
learns a BN (its parameters) with the structure passed in argument using a single pass estimation (no...
BNLearner< GUM_SCALAR > & EMenableEpsilon()
Enable the log-likelihood min diff stopping criterion in EM.
BNLearner< GUM_SCALAR > & setPCMaxCondSetSize(Size max_k)
BNLearner< GUM_SCALAR > & setPossibleEdges(const EdgeSet &set)
BNLearner< GUM_SCALAR > & allowArcDeletions(bool allow)
BNLearner< GUM_SCALAR > & setFCIExhaustiveSepSet(bool exhaustive)
BNLearner< GUM_SCALAR > & EMenableMaxTime()
enable EM's timeout stopping criterion
BayesNet< GUM_SCALAR > _learnParametersWithEM_(const DAG &dag, bool takeIntoAccountScore)
learns a BN (its parameters) with the structure passed in argument using the EM algorithm initialized...
BNLearner< GUM_SCALAR > & addNoChildrenNode(NodeId node)
BNLearner< GUM_SCALAR > & EMsetEpsilon(const double eps)
sets the stopping criterion of EM as being the minimal difference between two consecutive log-likelih...
BNLearner< GUM_SCALAR > & EMdisableMaxTime()
Disable EM's timeout stopping criterion.
BNLearner & operator=(const BNLearner &)
copy operator
bool greedyThickThinningReversals() const
BNLearner< GUM_SCALAR > & useGreedyHillClimbing()
BNLearner< GUM_SCALAR > & forbidEM()
prevent using the EM algorithm for parameter learning
BNLearner< GUM_SCALAR > & useK2(const Sequence< NodeId > &order)
BNLearner< GUM_SCALAR > & useExtendedGreedyHillClimbing()
BNLearner< GUM_SCALAR > & useScoreK2()
BNLearner< GUM_SCALAR > & setMandatoryArcs(const ArcSet &set)
BNLearner< GUM_SCALAR > & eraseNoChildrenNode(NodeId node)
BNLearner< GUM_SCALAR > & useMIIC()
bool fciExhaustiveSepSet() const
BNLearner< GUM_SCALAR > & EMsetPeriodSize(const Size p)
how many samples between 2 stoppings isEnabled
BayesNet< GUM_SCALAR > learnBN()
learn a Bayes Net from a file (must have read the db before)
BNLearner< GUM_SCALAR > & setForbiddenArcs(const ArcSet &set)
bool isConstraintBased() const
BNLearner< GUM_SCALAR > & useEMWithRateCriterion(const double epsilon, const double noise=default_EM_noise)
use The EM algorithm to learn parameters with the rate stopping criterion
BNLearner< GUM_SCALAR > & setInitialDAG(const DAG &dag)
BNLearner< GUM_SCALAR > & addNoParentNode(NodeId node)
BNLearner< GUM_SCALAR > & setPCAlpha(double alpha)
BNLearner< GUM_SCALAR > & useEMWithDiffCriterion(const double epsilon, const double noise=default_EM_noise)
use The EM algorithm to learn parameters with the diff stopping criterion
BNLearner< GUM_SCALAR > & useScoreLog2Likelihood()
BNLearner< GUM_SCALAR > & EMsetMaxIter(const Size max)
add a max iteration stopping criterion
BNLearner< GUM_SCALAR > & useEM(const double epsilon, const double noise=default_EM_noise)
use The EM algorithm to learn parameters
BNLearner< GUM_SCALAR > & useFCI()
bool isScoreBased() const
BNLearner< GUM_SCALAR > & useDirichletPrior(std::string_view filename, double weight=1)
BNLearner< GUM_SCALAR > & useLocalSearchWithTabuList(Size tabu_size=100, Size nb_decrease=2)
BNLearner< GUM_SCALAR > & EMsetMinEpsilonRate(const double rate)
sets the stopping criterion of EM as being the minimal log-likelihood's evolution rate
BNLearner< GUM_SCALAR > & allowArcTriangleDeletions(bool allow)
NodeProperty< Sequence< std::string > > _labelsFromBN_(std::string_view filename, const BayesNet< GUM_SCALAR > &src)
read the first line of a file to find column names
BNLearner< GUM_SCALAR > & useMDLCorrection()
BNLearner< GUM_SCALAR > & eraseNoParentNode(NodeId node)
BNLearner< GUM_SCALAR > & allowArcAdditions(bool allow)
BNLearner< GUM_SCALAR > & EMsetVerbosity(const bool v)
sets or unsets EM's verbosity
BNLearner< GUM_SCALAR > & eraseForbiddenArc(const Arc &arc)
BayesNet< GUM_SCALAR > learnParameters(const DAG &dag, bool takeIntoAccountScore=true)
learns a BN (its parameters) with the structure passed in argument
BNLearner< GUM_SCALAR > & setFCIAlpha(double alpha)
BNLearner< GUM_SCALAR > & useFCIG2Test()
void _checkDAGCompatibility_(const DAG &dag)
check that the database contains the nodes of the dag, else raise an exception
BNLearner< GUM_SCALAR > & useNMLCorrection()
BNLearner< GUM_SCALAR > & setSliceOrder(const NodeProperty< NodeId > &slice_order)
std::string toString() const
void createPrior_() override
create the prior used for learning
BNLearner< GUM_SCALAR > & EMsetMaxTime(const double timeout)
add a stopping criterion on timeout
void copyState(const BNLearner< GUM_SCALAR > &learner)
copy the states of the BNLearner
BNLearner< GUM_SCALAR > & setPCUnshieldedColliderSorted(bool sorted)
BNLearner< GUM_SCALAR > & EMenableMaxIter()
Enable stopping criterion on max iterations.
BNLearner< GUM_SCALAR > & useChi2Test()
BNLearner< GUM_SCALAR > & useScoreBDeu()
BNLearner< GUM_SCALAR > & setMaxIndegree(Size max_indegree)
Class for fast parsing of CSV file (never more than one line in application memory).
A DBRowGenerator class that returns the rows that are complete (fully observed) w....
A DBRowGenerator class that returns incomplete rows as EM would do.
the class used to read a row in the database and to transform it into a set of DBRow instances that c...
The class used to pack sets of generators.
void insertGenerator(const Generator &generator)
inserts a new generator at the end of the set
The base class for all DBRow generators.
DirichletPriorFromBN(const DatabaseTable &learning_db, const BayesNet< GUM_SCALAR > *priorbn)
default constructor
DirichletPriorFromDatabase(const DatabaseTable &learning_db, const DBRowGeneratorParser &prior_parser, const Bijection< NodeId, std::size_t > &nodeId2columns=Bijection< NodeId, std::size_t >())
default constructor
A pack of learning algorithms that can easily be used.
void usePC()
indicate that we wish to use PC (Chi2 test by default)
void eraseNoChildrenNode(NodeId node)
void EMenableEpsilon()
Enable the log-likelihood min diff stopping criterion in EM.
void EMsetPeriodSize(Size p)
how many samples between 2 stoppings isEnabled
void useGreedyHillClimbing()
indicate that we wish to use a greedy hill climbing algorithm
void useScoreBDeu()
indicate that we wish to use a BDeu score
void addNoParentNode(NodeId node)
void setSliceOrder(const NodeProperty< NodeId > &slice_order)
sets a partial order on the nodes
bool isScoreBased() const
indicate if the selected algorithm is score-based
void setForbiddenArcs(const ArcSet &set)
removes a total
void useFCIChi2Test()
indicate that we wish to use Chi2 independence test for FCI
void useBDeuPrior(double weight=1.0)
use the BDeu prior
void setMandatoryArcs(const ArcSet &set)
assign a set of mandatory arcs
bool greedyThickThinningReversals() const
returns whether arc reversals are allowed in the thin phase of greedy thick-thinning
void EMdisableMinEpsilonRate()
Disable the log-likelihood evolution rate stopping criterion.
void useExtendedGreedyHillClimbing()
indicate that we wish to use the extended greedy hill climbing algorithm
void setFCIMaxPathLength(Size max_len)
set maximum discriminating-path length for FCI R4 (default Size(-1) = unlimited)
void addMandatoryArc(const Arc &arc)
void EMenableMaxIter()
Enable stopping criterion on max iterations.
void useFCI()
indicate that we wish to use FCI (Chi2 test by default)
void useFCIG2Test()
indicate that we wish to use G2 independence test for FCI
void setMaxIndegree(Size max_indegree)
sets the max indegree
void EMdisableEpsilon()
Disable the min log-likelihood diff stopping criterion for EM.
void addPossibleEdge(const Edge &edge)
void EMsetMaxIter(Size max)
add a max iteration stopping criterion
void useChi2Test()
indicate that we wish to use Chi2 independence test for PC
void setInitialDAG(const DAG &)
sets an initial DAG structure
void useK2(const Sequence< NodeId > &order)
indicate that we wish to use K2
void allowArcDeletions(bool allow=true)
allow (true)/forbid (false) to delete arcs during learning.
void EMsetMinEpsilonRate(double rate)
sets the stopping criterion of EM as being the minimal log-likelihood's evolution rate
void setGreedyThickThinningReversals(bool allow)
enable or disable arc reversals in the thin phase of greedy thick-thinning
void EMdisableMaxIter()
Disable stopping criterion on max iterations.
void erasePossibleEdge(const Edge &edge)
void useScoreBIC()
indicate that we wish to use a BIC score
void allowArcTriangleDeletions(bool allow=true)
allow (true)/forbid (false) to delete arc triangles during learning.
void EMsetVerbosity(bool v)
sets or unsets EM's verbosity
void setPossibleEdges(const EdgeSet &set)
assign a set of possible edges
void useNoPrior()
use no prior
void eraseForbiddenArc(const Arc &arc)
void useSmoothingPrior(double weight=1)
use the prior smoothing
void allowArcAdditions(bool allow=true)
allow (true)/forbid (false) to add arcs during learning.
void useLocalSearchWithTabuList(Size tabu_size=100, Size nb_decrease=2)
indicate that we wish to use a local search with tabu list
void useScoreK2()
indicate that we wish to use a K2 score
void useGreedyThickThinning()
indicate that we wish to use greedy thick-thinning
void useG2Test()
indicate that we wish to use G2 independence test for PC
void setPossibleSkeleton(const UndiGraph &skeleton)
assign a set of possible edges
void useEMWithRateCriterion(const double epsilon, const double noise=default_EM_noise)
use The EM algorithm to learn parameters with the rate stopping criterion
void useNMLCorrection()
indicate that we wish to use the NML correction for and MIIC
void useEM(const double epsilon, const double noise=default_EM_noise)
use The EM algorithm to learn parameters
void useEMWithDiffCriterion(const double epsilon, const double noise=default_EM_noise)
use The EM algorithm to learn parameters with the diff stopping criterion
void forbidEM()
prevent using the EM algorithm for parameter learning
void setPCMaxCondSetSize(Size max_k)
set maximum conditioning set size for PC (default Size(-1) = unlimited)
void EMenableMaxTime()
sets the stopping criterion of EM as being the minimal difference between two consecutive log-likelih...
void useNoCorrection()
indicate that we wish to use the NoCorr correction for MIIC
void useScoreLog2Likelihood()
indicate that we wish to use a Log2Likelihood score
void useDirichletPrior(std::string_view filename, double weight=1)
use the Dirichlet prior from a database
void useMDLCorrection()
indicate that we wish to use the MDL correction for MIIC
void setFCIAlpha(double alpha)
set the significance threshold alpha for FCI (default 0.05)
bool fciExhaustiveSepSet() const
return true when FCI uses exhaustive sepset mode
void setPCAlpha(double alpha)
set the significance threshold alpha for PC (default 0.05)
void addForbiddenArc(const Arc &arc)
void addNoChildrenNode(NodeId node)
void EMsetMaxTime(double timeout)
add a stopping criterion on timeout
IBNLearner & operator=(const IBNLearner &)
copy operator
void EMenableMinEpsilonRate()
Enable the log-likelihood evolution rate stopping criterion.
void setFCIExhaustiveSepSet(bool exhaustive)
enable exhaustive sepset mode for FCI skeleton learning (default false)
void setTotalOrder(const Sequence< NodeId > &order)
sets a total order over some nodes
void useScoreAIC()
indicate that we wish to use an AIC score
void eraseMandatoryArc(const Arc &arc)
void allowArcReversals(bool allow=true)
allow (true)/forbid (false) to reverse arcs during learning.
void useMIIC()
indicate that we wish to use MIIC
void EMdisableMaxTime()
Disable EM's timeout stopping criterion.
void eraseNoParentNode(NodeId node)
void setPCUnshieldedColliderSorted(bool sorted)
set unshielded-collider ordering for PC: sorted=true uses descending p-value order (strongest evidenc...
void EMsetEpsilon(double eps)
sets the stopping criterion of EM as being the minimal difference between two consecutive log-likelih...
void setPCStable(bool stable)
set stable mode for PC — defer removals to end of each depth level (default true)
void useScoreBD()
indicate that we wish to use a BD score
bool isConstraintBased() const
indicate if the selected algorithm is constraint-based
NoPrior(const DatabaseTable &database, const Bijection< NodeId, std::size_t > &nodeId2columns=Bijection< NodeId, std::size_t >())
default constructor
the base class for all a priori
virtual void setWeight(double weight)
sets the weight of the a prior(kind of effective sample size)
SmoothingPrior(const DatabaseTable &database, const Bijection< NodeId, std::size_t > &nodeId2columns=Bijection< NodeId, std::size_t >())
default constructor
#define GUM_ERROR(type, msg)
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Size Idx
Type for indexes.
Set< Edge > EdgeSet
Some typdefs and define for shortcuts ...
Size NodeId
Type for node ids.
Set< Arc > ArcSet
Some typdefs and define for shortcuts ...
HashTable< NodeId, VAL > NodeProperty
Property on graph elements.
include the inlined functions if necessary
std::ostream & operator<<(std::ostream &stream, const IdCondSet &idset)
the display operator
gum is the global namespace for all aGrUM entities