52#ifndef GUM_LEARNING_GENERIC_BN_LEARNER_H
53#define GUM_LEARNING_GENERIC_BN_LEARNER_H
109 enum class ScoreType { AIC, BD, BDeu, BIC, fNML,
K2, LOG2LIKELIHOOD, MDL };
119 DIRICHLET_FROM_DATABASE,
120 DIRICHLET_FROM_BAYESNET,
127 GREEDY_HILL_CLIMBING,
128 LOCAL_SEARCH_WITH_TABU_LIST,
132 EXTENDED_GREEDY_HILL_CLIMBING,
133 GREEDY_THICK_THINNING
158 explicit Database(std::string_view file,
159 const std::vector< std::string >& missing_symbols,
160 const bool induceTypes =
false);
181 const std::vector< std::string >& missing_symbols);
190 template < GUM_Numeric GUM_SCALAR >
193 const std::vector< std::string >& missing_symbols);
228 const std::vector< std::size_t >&
domainSizes()
const;
231 const std::vector< std::string >&
names()
const;
253 std::size_t
nbRows()
const;
256 std::size_t
size()
const;
267 double weight(
const std::size_t i)
const;
296 template < GUM_Numeric GUM_SCALAR >
297 BayesNet< GUM_SCALAR > _BNVars_()
const;
318 const std::vector< std::string >& missingSymbols,
319 bool induceTypes =
true);
340 template < GUM_Numeric GUM_SCALAR >
343 const std::vector< std::string >& missing_symbols);
386 void setInitialDAG(
const DAG&);
392 const std::vector< std::string >& names()
const;
395 const std::vector< std::size_t >& domainSizes()
const;
397 Size domainSize(std::string_view var)
const;
404 NodeId idFromName(std::string_view var_name)
const;
411 void setDatabaseWeight(
const double new_weight);
417 void setRecordWeight(
const std::size_t i,
const double weight);
422 double recordWeight(
const std::size_t i)
const;
425 double databaseWeight()
const;
428 const std::string& nameFromId(
NodeId id)
const;
437 void useDatabaseRanges(
const std::vector< std::pair< std::size_t, std::size_t > >& new_ranges);
440 void clearDatabaseRanges();
446 const std::vector< std::pair< std::size_t, std::size_t > >& databaseRanges()
const;
469 std::pair< std::size_t, std::size_t > useCrossValidationFold(
const std::size_t learning_fold,
470 const std::size_t k_fold);
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 = {});
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 = {});
520 double logLikelihood(
const std::vector< NodeId >& vars,
521 const std::vector< NodeId >& knowing = {});
530 double logLikelihood(
const std::vector< std::string >& vars,
531 const std::vector< std::string >& knowing = {});
544 double mutualInformation(NodeId id1, NodeId id2,
const std::vector< NodeId >& knowing = {});
557 double mutualInformation(std::string_view var1,
558 std::string_view var2,
559 const std::vector< std::string >& knowing = {});
574 double correctedMutualInformation(NodeId id1,
576 const std::vector< NodeId >& knowing = {});
592 double correctedMutualInformation(std::string_view var1,
593 std::string_view var2,
594 const std::vector< std::string >& knowing = {});
604 double score(NodeId vars,
const std::vector< NodeId >& knowing = {});
615 double score(std::string_view vars,
const std::vector< std::string >& knowing = {});
622 std::vector< double > rawPseudoCount(
const std::vector< NodeId >& vars);
629 std::vector< double > rawPseudoCount(
const std::vector< std::string >& vars);
659 void useEM(
const double epsilon,
const double noise = default_EM_noise);
675 void useEMWithRateCriterion(
const double epsilon,
const double noise = default_EM_noise);
690 void useEMWithDiffCriterion(
const double epsilon,
const double noise = default_EM_noise);
696 bool isUsingEM()
const;
707 EMApproximationScheme& EM();
710 ApproximationSchemeSTATE EMState()
const;
713 std::string EMStateMessage()
const;
716 bool hasMissingValues()
const;
744 void useScoreLog2Likelihood();
763 void useBDeuPrior(
double weight = 1.0);
769 void useSmoothingPrior(
double weight = 1);
772 void useDirichletPrior(std::string_view filename,
double weight = 1);
777 std::string checkScorePriorCompatibility()
const;
786 void useGreedyHillClimbing();
801 void useExtendedGreedyHillClimbing();
807 void useLocalSearchWithTabuList(Size tabu_size = 100, Size nb_decrease = 2);
810 void useK2(
const Sequence< NodeId >& order);
813 void useK2(
const std::vector< NodeId >& order);
825 void useGreedyThickThinning();
828 void setGreedyThickThinningReversals(
bool allow);
831 bool greedyThickThinningReversals()
const;
834 bool isConstraintBased()
const;
837 bool isScoreBased()
const;
848 void allowArcAdditions(
bool allow =
true);
858 void allowArcDeletions(
bool allow =
true);
866 void allowArcReversals(
bool allow =
true);
873 void allowArcTriangleDeletions(
bool allow =
true);
883 void useNMLCorrection();
886 void useMDLCorrection();
889 void useNoCorrection();
893 std::vector< Arc > latentVariables()
const;
912 void setPCAlpha(
double alpha);
916 void setPCStable(
bool stable);
920 void setPCMaxCondSetSize(Size max_k);
925 void setPCUnshieldedColliderSorted(
bool sorted);
936 void useFCIChi2Test();
944 void setFCIAlpha(
double alpha);
948 void setFCIMaxPathLength(Size max_len);
952 void setFCIExhaustiveSepSet(
bool exhaustive);
956 bool fciExhaustiveSepSet()
const;
966 void setMaxIndegree(Size max_indegree);
973 void setSliceOrder(
const NodeProperty< NodeId >& slice_order);
981 void setSliceOrder(
const std::vector< std::vector< std::string > >& slices);
988 void unsetSliceOrder();
1007 void setTotalOrder(
const Sequence< NodeId >& order);
1008 void setTotalOrder(
const std::vector< std::string >& order);
1012 void unsetTotalOrder();
1017 void setForbiddenArcs(
const ArcSet& set);
1021 void addForbiddenArc(
const Arc& arc);
1022 void addForbiddenArc(NodeId tail, NodeId head);
1023 void addForbiddenArc(std::string_view tail, std::string_view head);
1028 void eraseForbiddenArc(
const Arc& arc);
1029 void eraseForbiddenArc(NodeId tail, NodeId head);
1030 void eraseForbiddenArc(std::string_view tail, std::string_view head);
1034 void setMandatoryArcs(
const ArcSet& set);
1038 void addMandatoryArc(
const Arc& arc);
1039 void addMandatoryArc(NodeId tail, NodeId head);
1040 void addMandatoryArc(std::string_view tail, std::string_view head);
1045 void eraseMandatoryArc(
const Arc& arc);
1046 void eraseMandatoryArc(NodeId tail, NodeId head);
1047 void eraseMandatoryArc(std::string_view tail, std::string_view head);
1052 void addNoParentNode(NodeId node);
1053 void addNoParentNode(std::string_view node);
1058 void eraseNoParentNode(NodeId node);
1059 void eraseNoParentNode(std::string_view node);
1063 void addNoChildrenNode(NodeId node);
1064 void addNoChildrenNode(std::string_view node);
1069 void eraseNoChildrenNode(NodeId node);
1070 void eraseNoChildrenNode(std::string_view node);
1077 void setPossibleEdges(
const EdgeSet& set);
1078 void setPossibleSkeleton(
const UndiGraph& skeleton);
1085 void addPossibleEdge(
const Edge& edge);
1086 void addPossibleEdge(NodeId tail, NodeId head);
1087 void addPossibleEdge(std::string_view tail, std::string_view head);
1092 void erasePossibleEdge(
const Edge& edge);
1093 void erasePossibleEdge(NodeId tail, NodeId head);
1094 void erasePossibleEdge(std::string_view tail, std::string_view head);
1118 void _setPriorWeight_(
double weight);
1256 std::vector< std::pair< std::size_t, std::size_t > >
ranges_;
1278 const std::vector< std::string >& missing_symbols);
1281 static void isCSVFileName_(std::string_view filename);
1291 bool take_into_account_score =
true);
1343 double epsilon()
const override;
1410 double maxTime()
const override;
1452 const std::vector< double >&
history()
const override;
1589 const std::vector< double >&
EMHistory()
const;
1598#ifndef GUM_NO_INLINE
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.
PC (Peter-Clark) constraint-based structure learning algorithm.
The SimpleMiic algorithm.
Class representing a Bayesian network.
Static math utilities for the chi2 distribution.
IApproximationSchemeConfiguration()
Class constructors.
ApproximationSchemeSTATE
The different state of an approximation scheme.
Base class for mixed graphs.
Partial Ancestral Graph: undirected topology with endpoint marks.
Base class for partially directed acyclic graphs.
ThreadNumberManager(Size nb_threads=0)
default constructor
A class that redirects gum_signal from algorithms to the listeners of BNLearn.
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.
The greedy hill climbing learning algorithm (for directed graphs).
The greedy thick-thinning learning algorithm (for directed graphs).
a helper to easily read databases
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
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)
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
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
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
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
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.
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
bool EMisEnabledMaxTime() const
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
AlgoType
an enumeration to select easily the learning algorithm to use
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
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
bool isEnabledMinEpsilonRate() const override
const std::vector< double > & EMHistory() const
returns the history of the last EM execution
Size nbDecreasingChanges_
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.
bool exhaustiveSepSetFci_
bool EMisEnabledMaxIter() const
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())
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 local search with tabu list learning algorithm (for directed graphs).
The MIIC learning algorithm.
the no a priorclass: corresponds to 0 weight-sample
PC (Peter-Clark) constraint-based structure learning algorithm.
The base class for estimating parameters of CPTs.
the base class for all a priori
The base class for all the scores used for learning (BIC, BDeu, etc).
The miic learning algorithm.
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.
std::size_t Size
In aGrUM, hashed values are unsigned long int.
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
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