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() {
101 GUM_DESTRUCTOR(BNLearner);
112 template < GUM_Numeric GUM_SCALAR >
113 BNLearner< GUM_SCALAR >&
114 BNLearner< GUM_SCALAR >::operator=(
const BNLearner< GUM_SCALAR >& src) {
115 IBNLearner::operator=(src);
120 template < GUM_Numeric GUM_SCALAR >
121 BNLearner< GUM_SCALAR >&
122 BNLearner< GUM_SCALAR >::operator=(BNLearner< GUM_SCALAR >&& src)
noexcept {
123 IBNLearner::operator=(std::move(src));
128 template < GUM_Numeric GUM_SCALAR >
129 BayesNet< GUM_SCALAR > BNLearner< GUM_SCALAR >::learnBN() {
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 >
144 void BNLearner< GUM_SCALAR >::_checkDAGCompatibility_(
const DAG& dag) {
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);
175 template < GUM_Numeric GUM_SCALAR >
176 BayesNet< GUM_SCALAR > BNLearner< GUM_SCALAR >::_learnParameters_(
const DAG& dag,
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()");
199 DBRowGeneratorParser parser(scoreDatabase_.databaseTable().handler(), DBRowGeneratorSet());
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 > >
218 BNLearner< GUM_SCALAR >::_initializeEMParameterLearning_(
const DAG& dag,
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(
238 DBRowGenerator4CompleteRows generator_bootstrap(col_types);
239 DBRowGeneratorSet genset_bootstrap;
240 genset_bootstrap.insertGenerator(generator_bootstrap);
241 DBRowGeneratorParser parser_bootstrap(database.handler(), genset_bootstrap);
242 std::shared_ptr< ParamEstimator > param_estimator_bootstrap(
243 createParamEstimator_(parser_bootstrap, takeIntoAccountScore));
246 BayesNet< GUM_SCALAR > dummy_bn;
247 DBRowGeneratorEM< GUM_SCALAR > generator_EM(col_types, dummy_bn);
248 DBRowGenerator& gen_EM = generator_EM;
249 DBRowGeneratorSet genset_EM;
250 genset_EM.insertGenerator(gen_EM);
251 DBRowGeneratorParser parser_EM(database.handler(), genset_EM);
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 >
265 BNLearner< GUM_SCALAR >::_learnParametersWithEM_(
const DAG& dag,
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 >
282 BNLearner< GUM_SCALAR >::_learnParametersWithEM_(
const BayesNet< GUM_SCALAR >& bn,
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 >
297 BayesNet< GUM_SCALAR > BNLearner< GUM_SCALAR >::learnParameters(
const DAG& dag,
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 >
311 BNLearner< GUM_SCALAR >::learnParameters(
const BayesNet< GUM_SCALAR >& bn,
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 >
333 BayesNet< GUM_SCALAR > BNLearner< GUM_SCALAR >::learnParameters(
bool take_into_account_score) {
334 return learnParameters(initialDag_, take_into_account_score);
337 template < GUM_Numeric GUM_SCALAR >
338 NodeProperty< Sequence< std::string > >
339 BNLearner< GUM_SCALAR >::_labelsFromBN_(std::string_view filename,
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();
351 NodeProperty< Sequence< std::string > > modals;
353 for (
gum::Idx col = 0; col < names.size(); col++) {
354 if (src.exists(names[col])) {
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 >
368 std::string BNLearner< GUM_SCALAR >::toString()
const {
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 > >
386 BNLearner< GUM_SCALAR >::state()
const {
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_) {
457 case CorrectedMutualInformation::KModeTypes::MDL :
458 vals.emplace_back(key,
"MDL",
"");
460 case CorrectedMutualInformation::KModeTypes::NML :
461 vals.emplace_back(key,
"NML",
"");
463 case CorrectedMutualInformation::KModeTypes::NoCorr :
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 >
605 void BNLearner< GUM_SCALAR >::copyState(
const BNLearner< GUM_SCALAR >& learner) {
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_) {
634 case CorrectedMutualInformation::KModeTypes::MDL : useMDLCorrection();
break;
635 case CorrectedMutualInformation::KModeTypes::NML : useNMLCorrection();
break;
636 case CorrectedMutualInformation::KModeTypes::NoCorr : useNoCorrection();
break;
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);
664 for (
const auto src: learner.constraintNoChildrenNodes_.nodes()) {
666 const auto dst = idFromName(learner.nameFromId(src));
667 addNoChildrenNode(dst);
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);
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);
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);
699 if (!learner.constraintSliceOrder_.sliceOrder().empty()) {
700 NodeProperty< NodeId > slice_order;
701 for (
const auto& p: learner.constraintSliceOrder_.sliceOrder()) {
703 slice_order.insert(idFromName(learner.nameFromId(p.first)), p.second);
708 setSliceOrder(slice_order);
710 if (!learner.constraintTotalOrder_.totalOrder().empty()) {
711 setTotalOrder(learner.constraintTotalOrder_.totalOrder());
715 template < GUM_Numeric GUM_SCALAR >
716 void BNLearner< GUM_SCALAR >::createPrior_() {
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());
740 prior_ =
new DirichletPriorFromDatabase(scoreDatabase_.databaseTable(),
741 priorDatabase_->parser(),
742 priorDatabase_->nodeId2Columns());
745 case BNLearnerPriorType::DIRICHLET_FROM_BAYESNET :
747 =
new DirichletPriorFromBN< GUM_SCALAR >(scoreDatabase_.databaseTable(), &_prior_bn_);
750 case BNLearnerPriorType::BDEU :
751 prior_ =
new BDeuPrior(scoreDatabase_.databaseTable(), scoreDatabase_.nodeId2Columns());
758 prior_->setWeight(priorWeight_);
761 if (old_prior !=
nullptr)
delete old_prior;
764 template < GUM_Numeric GUM_SCALAR >
765 std::ostream&
operator<<(std::ostream& output,
const BNLearner< GUM_SCALAR >& learner) {
766 output << learner.toString();
774 template < GUM_Numeric GUM_SCALAR >
775 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setInitialDAG(
const DAG& dag) {
776 IBNLearner::setInitialDAG(dag);
780 template < GUM_Numeric GUM_SCALAR >
781 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useEM(
const double epsilon,
782 const double noise) {
783 IBNLearner::useEM(epsilon, noise);
787 template < GUM_Numeric GUM_SCALAR >
788 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useEMWithRateCriterion(
const double epsilon,
789 const double noise) {
790 IBNLearner::useEMWithRateCriterion(epsilon, noise);
794 template < GUM_Numeric GUM_SCALAR >
795 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useEMWithDiffCriterion(
const double epsilon,
796 const double noise) {
797 IBNLearner::useEMWithDiffCriterion(epsilon, noise);
801 template < GUM_Numeric GUM_SCALAR >
802 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::forbidEM() {
803 IBNLearner::forbidEM();
807 template < GUM_Numeric GUM_SCALAR >
808 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMsetEpsilon(
const double eps) {
809 IBNLearner::EMsetEpsilon(eps);
813 template < GUM_Numeric GUM_SCALAR >
814 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMdisableEpsilon() {
815 IBNLearner::EMdisableEpsilon();
819 template < GUM_Numeric GUM_SCALAR >
820 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMenableEpsilon() {
821 IBNLearner::EMenableEpsilon();
825 template < GUM_Numeric GUM_SCALAR >
826 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMsetMinEpsilonRate(
const double rate) {
827 IBNLearner::EMsetMinEpsilonRate(rate);
831 template < GUM_Numeric GUM_SCALAR >
832 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMdisableMinEpsilonRate() {
833 IBNLearner::EMdisableMinEpsilonRate();
837 template < GUM_Numeric GUM_SCALAR >
838 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMenableMinEpsilonRate() {
839 IBNLearner::EMenableMinEpsilonRate();
843 template < GUM_Numeric GUM_SCALAR >
844 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMsetMaxIter(
const Size max) {
845 IBNLearner::EMsetMaxIter(max);
849 template < GUM_Numeric GUM_SCALAR >
850 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMdisableMaxIter() {
851 IBNLearner::EMdisableMaxIter();
855 template < GUM_Numeric GUM_SCALAR >
856 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMenableMaxIter() {
857 IBNLearner::EMenableMaxIter();
861 template < GUM_Numeric GUM_SCALAR >
862 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMsetMaxTime(
const double timeout) {
863 IBNLearner::EMsetMaxTime(timeout);
867 template < GUM_Numeric GUM_SCALAR >
868 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMdisableMaxTime() {
869 IBNLearner::EMdisableMaxTime();
873 template < GUM_Numeric GUM_SCALAR >
874 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMenableMaxTime() {
875 IBNLearner::EMenableMaxTime();
879 template < GUM_Numeric GUM_SCALAR >
880 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMsetPeriodSize(
const Size p) {
881 IBNLearner::EMsetPeriodSize(p);
885 template < GUM_Numeric GUM_SCALAR >
886 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::EMsetVerbosity(
const bool v) {
887 IBNLearner::EMsetVerbosity(v);
891 template < GUM_Numeric GUM_SCALAR >
892 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useScoreAIC() {
893 IBNLearner::useScoreAIC();
897 template < GUM_Numeric GUM_SCALAR >
898 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useScoreBD() {
899 IBNLearner::useScoreBD();
903 template < GUM_Numeric GUM_SCALAR >
904 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useScoreBDeu() {
905 IBNLearner::useScoreBDeu();
909 template < GUM_Numeric GUM_SCALAR >
910 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useScoreBIC() {
911 IBNLearner::useScoreBIC();
915 template < GUM_Numeric GUM_SCALAR >
916 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useScoreK2() {
917 IBNLearner::useScoreK2();
921 template < GUM_Numeric GUM_SCALAR >
922 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useScoreLog2Likelihood() {
923 IBNLearner::useScoreLog2Likelihood();
927 template < GUM_Numeric GUM_SCALAR >
928 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useNoPrior() {
929 IBNLearner::useNoPrior();
933 template < GUM_Numeric GUM_SCALAR >
934 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useBDeuPrior(
double weight) {
935 IBNLearner::useBDeuPrior(weight);
939 template < GUM_Numeric GUM_SCALAR >
940 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useSmoothingPrior(
double weight) {
941 IBNLearner::useSmoothingPrior(weight);
945 template < GUM_Numeric GUM_SCALAR >
946 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useDirichletPrior(std::string_view filename,
948 IBNLearner::useDirichletPrior(filename, weight);
952 template < GUM_Numeric GUM_SCALAR >
953 BNLearner< GUM_SCALAR >&
957 priorType_ = BNLearnerPriorType::DIRICHLET_FROM_BAYESNET;
958 _setPriorWeight_(weight);
962 template < GUM_Numeric GUM_SCALAR >
963 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useGreedyHillClimbing() {
964 IBNLearner::useGreedyHillClimbing();
968 template < GUM_Numeric GUM_SCALAR >
969 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useExtendedGreedyHillClimbing() {
970 IBNLearner::useExtendedGreedyHillClimbing();
974 template < GUM_Numeric GUM_SCALAR >
975 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useGreedyThickThinning() {
976 IBNLearner::useGreedyThickThinning();
980 template < GUM_Numeric GUM_SCALAR >
981 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setGreedyThickThinningReversals(
bool allow) {
982 IBNLearner::setGreedyThickThinningReversals(allow);
986 template < GUM_Numeric GUM_SCALAR >
987 bool BNLearner< GUM_SCALAR >::greedyThickThinningReversals()
const {
988 return IBNLearner::greedyThickThinningReversals();
991 template < GUM_Numeric GUM_SCALAR >
992 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useLocalSearchWithTabuList(Size tabu_size,
994 IBNLearner::useLocalSearchWithTabuList(tabu_size, nb_decrease);
998 template < GUM_Numeric GUM_SCALAR >
999 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useK2(
const Sequence< NodeId >& order) {
1000 IBNLearner::useK2(order);
1004 template < GUM_Numeric GUM_SCALAR >
1005 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useK2(
const std::vector< NodeId >& order) {
1006 IBNLearner::useK2(order);
1010 template < GUM_Numeric GUM_SCALAR >
1011 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useMIIC() {
1012 IBNLearner::useMIIC();
1016 template < GUM_Numeric GUM_SCALAR >
1017 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::usePC() {
1018 IBNLearner::usePC();
1022 template < GUM_Numeric GUM_SCALAR >
1023 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useFCI() {
1024 IBNLearner::useFCI();
1028 template < GUM_Numeric GUM_SCALAR >
1029 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useFCIChi2Test() {
1030 IBNLearner::useFCIChi2Test();
1034 template < GUM_Numeric GUM_SCALAR >
1035 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useFCIG2Test() {
1036 IBNLearner::useFCIG2Test();
1040 template < GUM_Numeric GUM_SCALAR >
1041 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setFCIAlpha(
double alpha) {
1042 IBNLearner::setFCIAlpha(alpha);
1046 template < GUM_Numeric GUM_SCALAR >
1047 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setFCIMaxPathLength(Size max_len) {
1048 IBNLearner::setFCIMaxPathLength(max_len);
1052 template < GUM_Numeric GUM_SCALAR >
1053 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setFCIExhaustiveSepSet(
bool exhaustive) {
1054 IBNLearner::setFCIExhaustiveSepSet(exhaustive);
1058 template < GUM_Numeric GUM_SCALAR >
1059 bool BNLearner< GUM_SCALAR >::fciExhaustiveSepSet()
const {
1060 return IBNLearner::fciExhaustiveSepSet();
1063 template < GUM_Numeric GUM_SCALAR >
1064 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useChi2Test() {
1065 IBNLearner::useChi2Test();
1069 template < GUM_Numeric GUM_SCALAR >
1070 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useG2Test() {
1071 IBNLearner::useG2Test();
1075 template < GUM_Numeric GUM_SCALAR >
1076 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setPCAlpha(
double alpha) {
1077 IBNLearner::setPCAlpha(alpha);
1081 template < GUM_Numeric GUM_SCALAR >
1082 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setPCStable(
bool stable) {
1083 IBNLearner::setPCStable(stable);
1087 template < GUM_Numeric GUM_SCALAR >
1088 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setPCMaxCondSetSize(Size max_k) {
1089 IBNLearner::setPCMaxCondSetSize(max_k);
1093 template < GUM_Numeric GUM_SCALAR >
1094 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setPCUnshieldedColliderSorted(
bool sorted) {
1095 IBNLearner::setPCUnshieldedColliderSorted(sorted);
1099 template < GUM_Numeric GUM_SCALAR >
1100 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useNMLCorrection() {
1101 IBNLearner::useNMLCorrection();
1105 template < GUM_Numeric GUM_SCALAR >
1106 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useMDLCorrection() {
1107 IBNLearner::useMDLCorrection();
1111 template < GUM_Numeric GUM_SCALAR >
1112 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::useNoCorrection() {
1113 IBNLearner::useNoCorrection();
1117 template < GUM_Numeric GUM_SCALAR >
1118 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setMaxIndegree(Size max_indegree) {
1119 IBNLearner::setMaxIndegree(max_indegree);
1123 template < GUM_Numeric GUM_SCALAR >
1124 BNLearner< GUM_SCALAR >&
1125 BNLearner< GUM_SCALAR >::setSliceOrder(
const NodeProperty< NodeId >& slice_order) {
1126 IBNLearner::setSliceOrder(slice_order);
1130 template < GUM_Numeric GUM_SCALAR >
1131 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setSliceOrder(
1132 const std::vector< std::vector< std::string > >& slices) {
1133 IBNLearner::setSliceOrder(slices);
1137 template < GUM_Numeric GUM_SCALAR >
1138 BNLearner< GUM_SCALAR >&
1139 BNLearner< GUM_SCALAR >::setTotalOrder(
const std::vector< std::string >& order) {
1140 IBNLearner::setTotalOrder(order);
1144 template < GUM_Numeric GUM_SCALAR >
1145 BNLearner< GUM_SCALAR >&
1146 BNLearner< GUM_SCALAR >::setTotalOrder(
const Sequence< NodeId >& order) {
1147 IBNLearner::setTotalOrder(order);
1151 template < GUM_Numeric GUM_SCALAR >
1152 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setForbiddenArcs(
const ArcSet& set) {
1153 IBNLearner::setForbiddenArcs(set);
1157 template < GUM_Numeric GUM_SCALAR >
1158 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addForbiddenArc(
const Arc& arc) {
1159 IBNLearner::addForbiddenArc(arc);
1163 template < GUM_Numeric GUM_SCALAR >
1164 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addForbiddenArc(NodeId tail, NodeId head) {
1165 IBNLearner::addForbiddenArc(tail, head);
1169 template < GUM_Numeric GUM_SCALAR >
1170 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addForbiddenArc(std::string_view tail,
1171 std::string_view head) {
1172 IBNLearner::addForbiddenArc(tail, head);
1176 template < GUM_Numeric GUM_SCALAR >
1177 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseForbiddenArc(
const Arc& arc) {
1178 IBNLearner::eraseForbiddenArc(arc);
1182 template < GUM_Numeric GUM_SCALAR >
1183 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseForbiddenArc(NodeId tail, NodeId head) {
1184 IBNLearner::eraseForbiddenArc(tail, head);
1188 template < GUM_Numeric GUM_SCALAR >
1189 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseForbiddenArc(std::string_view tail,
1190 std::string_view head) {
1191 IBNLearner::eraseForbiddenArc(tail, head);
1195 template < GUM_Numeric GUM_SCALAR >
1196 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addMandatoryArc(
const Arc& arc) {
1197 IBNLearner::addMandatoryArc(arc);
1201 template < GUM_Numeric GUM_SCALAR >
1202 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addMandatoryArc(NodeId tail, NodeId head) {
1203 IBNLearner::addMandatoryArc(tail, head);
1207 template < GUM_Numeric GUM_SCALAR >
1208 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addMandatoryArc(std::string_view tail,
1209 std::string_view head) {
1210 IBNLearner::addMandatoryArc(tail, head);
1214 template < GUM_Numeric GUM_SCALAR >
1215 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseMandatoryArc(
const Arc& arc) {
1216 IBNLearner::eraseMandatoryArc(arc);
1220 template < GUM_Numeric GUM_SCALAR >
1221 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseMandatoryArc(NodeId tail, NodeId head) {
1222 IBNLearner::eraseMandatoryArc(tail, head);
1226 template < GUM_Numeric GUM_SCALAR >
1227 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseMandatoryArc(std::string_view tail,
1228 std::string_view head) {
1229 IBNLearner::eraseMandatoryArc(tail, head);
1233 template < GUM_Numeric GUM_SCALAR >
1234 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addPossibleEdge(
const Edge& edge) {
1235 IBNLearner::addPossibleEdge(edge);
1239 template < GUM_Numeric GUM_SCALAR >
1240 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addPossibleEdge(NodeId tail, NodeId head) {
1241 IBNLearner::addPossibleEdge(tail, head);
1245 template < GUM_Numeric GUM_SCALAR >
1246 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addPossibleEdge(std::string_view tail,
1247 std::string_view head) {
1248 IBNLearner::addPossibleEdge(tail, head);
1252 template < GUM_Numeric GUM_SCALAR >
1253 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::erasePossibleEdge(
const Edge& edge) {
1254 IBNLearner::erasePossibleEdge(edge);
1258 template < GUM_Numeric GUM_SCALAR >
1259 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::erasePossibleEdge(NodeId tail, NodeId head) {
1260 IBNLearner::erasePossibleEdge(tail, head);
1264 template < GUM_Numeric GUM_SCALAR >
1265 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::erasePossibleEdge(std::string_view tail,
1266 std::string_view head) {
1267 IBNLearner::erasePossibleEdge(tail, head);
1271 template < GUM_Numeric GUM_SCALAR >
1272 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setMandatoryArcs(
const ArcSet& set) {
1273 IBNLearner::setMandatoryArcs(set);
1277 template < GUM_Numeric GUM_SCALAR >
1278 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::setPossibleEdges(
const EdgeSet& set) {
1279 IBNLearner::setPossibleEdges(set);
1283 template < GUM_Numeric GUM_SCALAR >
1284 BNLearner< GUM_SCALAR >&
1285 BNLearner< GUM_SCALAR >::setPossibleSkeleton(
const UndiGraph& skeleton) {
1286 IBNLearner::setPossibleSkeleton(skeleton);
1290 template < GUM_Numeric GUM_SCALAR >
1291 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addNoParentNode(NodeId node) {
1292 IBNLearner::addNoParentNode(node);
1296 template < GUM_Numeric GUM_SCALAR >
1297 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addNoParentNode(std::string_view name) {
1298 IBNLearner::addNoParentNode(name);
1302 template < GUM_Numeric GUM_SCALAR >
1303 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseNoParentNode(NodeId node) {
1304 IBNLearner::eraseNoParentNode(node);
1308 template < GUM_Numeric GUM_SCALAR >
1309 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseNoParentNode(std::string_view name) {
1310 IBNLearner::eraseNoParentNode(name);
1314 template < GUM_Numeric GUM_SCALAR >
1315 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addNoChildrenNode(NodeId node) {
1316 IBNLearner::addNoChildrenNode(node);
1320 template < GUM_Numeric GUM_SCALAR >
1321 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::addNoChildrenNode(std::string_view name) {
1322 IBNLearner::addNoChildrenNode(name);
1326 template < GUM_Numeric GUM_SCALAR >
1327 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseNoChildrenNode(NodeId node) {
1328 IBNLearner::eraseNoChildrenNode(node);
1332 template < GUM_Numeric GUM_SCALAR >
1333 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::eraseNoChildrenNode(std::string_view name) {
1334 IBNLearner::eraseNoChildrenNode(name);
1338 template < GUM_Numeric GUM_SCALAR >
1339 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::allowArcAdditions(
bool allow) {
1340 IBNLearner::allowArcAdditions(allow);
1344 template < GUM_Numeric GUM_SCALAR >
1345 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::allowArcDeletions(
bool allow) {
1346 IBNLearner::allowArcDeletions(allow);
1350 template < GUM_Numeric GUM_SCALAR >
1351 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::allowArcReversals(
bool allow) {
1352 IBNLearner::allowArcReversals(allow);
1356 template < GUM_Numeric GUM_SCALAR >
1357 BNLearner< GUM_SCALAR >& BNLearner< GUM_SCALAR >::allowArcTriangleDeletions(
bool allow) {
1358 IBNLearner::allowArcTriangleDeletions(allow);
1362 template < GUM_Numeric GUM_SCALAR >
1363 bool BNLearner< GUM_SCALAR >::isConstraintBased()
const {
1364 return IBNLearner::isConstraintBased();
1367 template < GUM_Numeric GUM_SCALAR >
1368 bool BNLearner< GUM_SCALAR >::isScoreBased()
const {
1369 return IBNLearner::isScoreBased();
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.
Error: The database contains some missing values.
Error: A name of variable is not found in the database.
Exception : operation not allowed.
BNLearner(std::string_view filename, const std::vector< std::string > &missingSymbols={"?"}, const bool induceTypes=true)
default constructor
A pack of learning algorithms that can easily be used.
#define GUM_ERROR(type, msg)
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Size Idx
Type for indexes.
Size NodeId
Type for node ids.
include the inlined functions if necessary
gum is the global namespace for all aGrUM entities
std::ostream & operator<<(std::ostream &out, const TiXmlNode &base)