65 template < GUM_Numeric GUM_SCALAR >
73 template < GUM_Numeric GUM_SCALAR >
76 std::string_view head,
78 const int last =
static_cast< int >(
_prior_ktbn_.k()) - 1;
93 template < GUM_Numeric GUM_SCALAR >
96 std::string_view headBase,
100 "unknown base variable '" << tailBase
101 <<
"': it is not one of this learner's variables")
104 "unknown base variable '" << headBase
105 <<
"': it is not one of this learner's variables")
109 const bool tailAtemp =
_prior_ktbn_.atemporalVarNames().contains(std::string{tailBase});
110 const bool headAtemp =
_prior_ktbn_.atemporalVarNames().contains(std::string{headBase});
112 if (tailAtemp && headAtemp) {
114 }
else if (tailAtemp) {
115 for (
int hs = 0; hs <
k; ++hs)
117 }
else if (headAtemp) {
120 for (
int ts = 0; ts <
k; ++ts)
121 for (
int hs = ts; hs <
k; ++hs)
130 template < GUM_Numeric GUM_SCALAR >
132 std::string_view csvBaseName,
135 const std::unordered_set< std::string >& atemporalVars,
136 const std::vector< std::string >& missingSymbols,
138 bool ignoreMissingSymbols) :
150 }
catch (
const gum::UnknownLabelInDatabase&) {
157 "KTBNLearner CSV constructor: an unknown label was encountered while "
158 "reading the trajectory CSVs. The variable domains are inferred from "
159 "the first CSV alone, so any variable whose modalities are not all "
160 "present in trajectory 1 will trigger this error (atemporal variables "
161 "are especially prone: each trajectory holds a single constant value "
162 "for them, so at most one label appears in trajectory 1). "
163 "Use the BN-schema constructor "
164 "KTBNLearner(dir, base, n, k, bn, atemporals) to supply the full "
165 "variable domains explicitly.")
172 template < GUM_Numeric GUM_SCALAR >
174 std::string_view csvBaseName,
177 const std::vector< std::string >& missingSymbols,
179 bool ignoreMissingSymbols) :
195 ignoreMissingSymbols) {}
197 template < GUM_Numeric GUM_SCALAR >
199 std::string_view csvBaseName,
202 const BayesNet< GUM_SCALAR >& bn,
203 const std::unordered_set< std::string >& atemporalVars,
204 const std::vector< std::string >& missingSymbols,
205 bool ignoreMissingSymbols) :
219 template < GUM_Numeric GUM_SCALAR >
228 template < GUM_Numeric GUM_SCALAR >
235 const bool suppressAtemporal
241 : BayesNet< GUM_SCALAR >{};
242 return _assemble_(transitionBN, initialBN, atemporalBN);
245 template < GUM_Numeric GUM_SCALAR >
247 bool takeIntoAccountScore) {
250 "learnParameters: structure has k="
251 << structure.k() <<
" but this learner was built with k=" <<
_prior_ktbn_.k())
265 for (std::size_t
id = 0;
id < nbTransNodes; ++id)
269 for (std::size_t
id = 0;
id < nbInitNodes; ++id)
274 for (std::size_t
id = 0;
id < nbAtemNodes; ++id)
278 for (
const auto& [tail, head]: structure.arcs()) {
279 const auto& [tailBase, tailSlice] = tail;
280 const auto& [headBase, headSlice] = head;
281 const std::string tailName =
_encode_(tailBase, tailSlice);
282 const std::string headName =
_encode_(headBase, headSlice);
284 if (headSlice ==
k - 1) {
297 BayesNet< GUM_SCALAR > transitionBN
299 BayesNet< GUM_SCALAR > initialBN
301 BayesNet< GUM_SCALAR > atemporalBN
304 : BayesNet< GUM_SCALAR >{};
306 return _assemble_(transitionBN, initialBN, atemporalBN);
314 template < GUM_Numeric GUM_SCALAR >
320 template < GUM_Numeric GUM_SCALAR >
326 template < GUM_Numeric GUM_SCALAR >
332 template < GUM_Numeric GUM_SCALAR >
338 template < GUM_Numeric GUM_SCALAR >
340 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreLog2Likelihood(); });
344 template < GUM_Numeric GUM_SCALAR >
350 template < GUM_Numeric GUM_SCALAR >
355 template < GUM_Numeric GUM_SCALAR >
367 template < GUM_Numeric GUM_SCALAR >
369 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useGreedyHillClimbing(); });
373 template < GUM_Numeric GUM_SCALAR >
375 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useExtendedGreedyHillClimbing(); });
379 template < GUM_Numeric GUM_SCALAR >
383 [&](BNLearner< GUM_SCALAR >& l) { l.useLocalSearchWithTabuList(tabu_size, nb_decrease); });
387 template < GUM_Numeric GUM_SCALAR >
397 template < GUM_Numeric GUM_SCALAR >
403 template < GUM_Numeric GUM_SCALAR >
409 template < GUM_Numeric GUM_SCALAR >
415 template < GUM_Numeric GUM_SCALAR >
416 std::vector< std::pair< std::string, std::string > >
423 std::vector< std::pair< std::string, std::string > > result;
424 std::unordered_set< std::string > seen;
426 auto collect = [&](
const std::unique_ptr< BNLearner< GUM_SCALAR > >& learner) {
427 if (!learner)
return;
428 for (
const auto& arc: learner->latentVariables()) {
429 std::string tail = learner->nameFromId(arc.tail());
430 std::string head = learner->nameFromId(arc.head());
431 if (seen.insert(tail +
'\t' + head).second)
432 result.emplace_back(std::move(tail), std::move(head));
445 template < GUM_Numeric GUM_SCALAR >
447 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.useSmoothingPrior(weight); });
455 template < GUM_Numeric GUM_SCALAR >
457 std::string_view headNode) {
459 l.addForbiddenArc(tailNode, headNode);
464 template < GUM_Numeric GUM_SCALAR >
467 std::string_view headBase,
473 template < GUM_Numeric GUM_SCALAR >
476 std::string_view headNode) {
481 l.eraseForbiddenArc(tailNode, headNode);
486 template < GUM_Numeric GUM_SCALAR >
489 std::string_view headBase,
494 template < GUM_Numeric GUM_SCALAR >
496 std::string_view headNode) {
505 l.addMandatoryArc(tailNode, headNode);
510 template < GUM_Numeric GUM_SCALAR >
513 std::string_view headBase,
519 template < GUM_Numeric GUM_SCALAR >
522 std::string_view headNode) {
526 l.eraseMandatoryArc(tailNode, headNode);
531 template < GUM_Numeric GUM_SCALAR >
534 std::string_view headBase,
539 template < GUM_Numeric GUM_SCALAR >
542 std::string_view headBase) {
548 for (
int t = 0; t <
k; ++t)
553 template < GUM_Numeric GUM_SCALAR >
556 std::string_view headBase) {
560 for (
int t = 0; t <
k; ++t)
565 template < GUM_Numeric GUM_SCALAR >
568 std::string_view headBase) {
575 template < GUM_Numeric GUM_SCALAR >
578 std::string_view headBase) {
585 template < GUM_Numeric GUM_SCALAR >
591 template < GUM_Numeric GUM_SCALAR >
599 }
else if (!isAtemp) {
606 template < GUM_Numeric GUM_SCALAR >
612 template < GUM_Numeric GUM_SCALAR >
626 template < GUM_Numeric GUM_SCALAR >
632 template < GUM_Numeric GUM_SCALAR >
642 template < GUM_Numeric GUM_SCALAR >
648 template < GUM_Numeric GUM_SCALAR >
658 template < GUM_Numeric GUM_SCALAR >
661 std::string_view headBase,
668 template < GUM_Numeric GUM_SCALAR >
671 std::string_view headBase,
676 template < GUM_Numeric GUM_SCALAR >
678 std::string_view head) {
691 if (tailSlice < last && headSlice < last)
_initialLearner_->addPossibleEdge(tail, head);
702 template < GUM_Numeric GUM_SCALAR >
704 std::string_view head) {
710 if (tailSlice < last && headSlice < last)
_initialLearner_->erasePossibleEdge(tail, head);
721 template < GUM_Numeric GUM_SCALAR >
723 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.allowArcAdditions(allow); });
727 template < GUM_Numeric GUM_SCALAR >
729 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.allowArcDeletions(allow); });
733 template < GUM_Numeric GUM_SCALAR >
735 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.allowArcReversals(allow); });
739 template < GUM_Numeric GUM_SCALAR >
741 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.setMaxIndegree(max_indegree); });
749 template < GUM_Numeric GUM_SCALAR >
754 template < GUM_Numeric GUM_SCALAR >
759 template < GUM_Numeric GUM_SCALAR >
767 template < GUM_Numeric GUM_SCALAR >
772 template < GUM_Numeric GUM_SCALAR >
777 template < GUM_Numeric GUM_SCALAR >
782 s <<
"k : " <<
k() <<
'\n';
784 <<
" temporal, " <<
_prior_ktbn_.nbAtemporalVars() <<
" atemporal)" <<
'\n';
788 s <<
"=== Transition learner (arcs into slice " << (
k() - 1) <<
") ===" <<
'\n';
791 s <<
"=== Initial learner (slices 0.." << (
k() - 2) <<
") ===" <<
'\n';
795 s <<
"=== Atemporal learner (atemporal->atemporal arcs) ===" <<
'\n';
801 template < GUM_Numeric GUM_SCALAR >
802 std::vector< std::tuple< std::string, std::string, std::string > >
812 for (
auto& [key, val, comment]: result) {
813 if (key !=
"Variables")
continue;
815 std::string collapsed;
817 auto emit = [&](
const std::string& base,
int slice) {
819 if (!first) collapsed +=
", ";
820 collapsed += base +
"[" + std::to_string(var.domainSize()) +
"]";
825 for (
const auto& base:
_prior_ktbn_.atemporalVarNames())
833 template < GUM_Numeric GUM_SCALAR >
855 template < GUM_Numeric GUM_SCALAR >
862 template < GUM_Numeric GUM_SCALAR >
867 template < GUM_Numeric GUM_SCALAR >
872 template < GUM_Numeric GUM_SCALAR >
878 template < GUM_Numeric GUM_SCALAR >
886 std::vector< std::string > result;
888 for (
Size i = 0; i < n; ++i)
893 template < GUM_Numeric GUM_SCALAR >
898 return std::vector< Size >(all.begin(), all.begin() +
nbCols());
901 template < GUM_Numeric GUM_SCALAR >
907 const std::string b{base};
918 template < GUM_Numeric GUM_SCALAR >
920 std::string_view dirPath,
921 std::string_view csvBaseName,
924 const std::vector< std::string >& missingSymbols) {
932 template < GUM_Numeric GUM_SCALAR >
934 std::string_view dirPath,
935 std::string_view csvBaseName,
937 const std::unordered_set< std::string >& atemporalVars,
938 const std::vector< std::string >& missingSymbols,
942 namespace fs = std::filesystem;
943 const std::string firstCSV
944 = (fs::path{dirPath} / (std::string{csvBaseName} +
"1.csv")).
string();
952 const BNLearner< GUM_SCALAR > tmpLearner(firstCSV, missingSymbols, induceTypes);
956 const DBTranslatorSet& translators = tmpLearner.database().translatorSet();
957 const std::vector< std::string >&
names = tmpLearner.names();
959 KTBN< GUM_SCALAR > prior(
k);
960 for (std::size_t i = 0; i <
names.size(); ++i) {
962 !atemporalVars.contains(
names[i]));
966 for (
const std::string& aname: atemporalVars)
967 if (!prior.exists(aname))
969 "atemporal variable '" << aname <<
"' not found in the CSV header")
973 template < GUM_Numeric GUM_SCALAR >
976 const BayesNet< GUM_SCALAR >& bn,
977 const std::unordered_set< std::string >& atemporalVars) {
980 KTBN< GUM_SCALAR > prior(
k);
983 for (
const NodeId node: bn.nodes()) {
985 prior.add(var, !atemporalVars.contains(var.
name()));
987 for (
const std::string& aname: atemporalVars)
988 if (!prior.exists(aname))
993 template < GUM_Numeric GUM_SCALAR >
995 std::string_view csvBaseName,
997 const std::vector< std::string >& missingSymbols) {
1023 const std::unordered_set< std::string > missingSet(missingSymbols.begin(),
1024 missingSymbols.end());
1025 const auto insertIfComplete = [&](
DatabaseTable& table,
const std::vector< std::string >& r) {
1027 for (
const auto& cell: r)
1028 if (missingSet.contains(cell)) {
1036 const std::filesystem::path dir{dirPath};
1037 const std::string stem{csvBaseName};
1039 const std::size_t transRowSize = nbAtempVars + nbTempVars *
k;
1040 const std::size_t initRowSize = nbAtempVars + nbTempVars * (
k - 1);
1041 const std::size_t atempRowSize = nbAtempVars;
1043 std::vector< std::string > header;
1044 std::unordered_set< Size > atemVarsCols;
1047 std::vector< std::vector< std::string > > buffer;
1048 std::vector< std::string > row;
1049 row.reserve(transRowSize);
1054 const std::filesystem::path file = dir / (stem + std::to_string(i + 1) +
".csv");
1055 std::ifstream is(file, std::ifstream::in);
1056 if (!is.is_open())
GUM_ERROR(gum::IOError,
"Cannot open " << file.string());
1065 const auto& rawHeader = parser.
current();
1066 header.assign(rawHeader.begin(), rawHeader.end());
1067 for (std::size_t c = 0; c < header.size(); ++c)
1068 if (
_prior_ktbn_.atemporalVarNames().contains(header[c])) atemVarsCols.insert(c);
1074 const std::unordered_set< std::string > headerSet(header.begin(), header.end());
1075 auto requirePresent = [&](
const std::string& base) {
1076 if (!headerSet.contains(base))
1078 "schema variable '" << base <<
"' is absent from '" << file.string() <<
"'")
1080 for (
const auto& base:
_prior_ktbn_.temporalVarNames())
1081 requirePresent(base);
1082 for (
const auto& base:
_prior_ktbn_.atemporalVarNames())
1083 requirePresent(base);
1088 for (
const std::string& col: header)
1091 "CSV column '" << col <<
"' in '" << file.string()
1092 <<
"' is not declared as a variable of this KTBNLearner")
1098 std::vector< std::string > varNamesTran;
1099 std::vector< std::string > varNamesInit;
1100 std::vector< std::string > varNamesAtemp;
1103 = [&](
DatabaseTable& table,
const std::string& base,
int slice, std::size_t col) {
1110 varNamesTran.reserve(transRowSize);
1111 varNamesInit.reserve(initRowSize);
1112 varNamesAtemp.reserve(atempRowSize);
1114 std::size_t tcol = 0, icol = 0, acol = 0;
1115 for (std::size_t c = 0; c < header.size(); ++c) {
1117 insertTrans(transitionTable, header[c], slice, tcol++);
1118 varNamesTran.push_back(
_encode_(header[c], slice));
1119 insertTrans(initTable, header[c], slice, icol++);
1120 varNamesInit.push_back(
_encode_(header[c], slice));
1121 if (atemVarsCols.contains(c)) {
1123 varNamesAtemp.push_back(header[c]);
1126 for (
Size slice = 1; slice <
k - 1; ++slice)
1127 for (std::size_t c = 0; c < header.size(); ++c)
1128 if (!atemVarsCols.contains(c)) {
1129 insertTrans(initTable, header[c], (
int)slice, icol++);
1130 varNamesInit.push_back(
_encode_(header[c], (
int)slice));
1132 for (
Size slice = 1; slice <
k; ++slice)
1133 for (std::size_t c = 0; c < header.size(); ++c)
1134 if (!atemVarsCols.contains(c)) {
1135 insertTrans(transitionTable, header[c], (
int)slice, tcol++);
1136 varNamesTran.push_back(
_encode_(header[c], (
int)slice));
1145 const auto& raw = parser.
current();
1146 bool same = (raw.size() == header.size());
1147 for (std::size_t c = 0; same && c < header.size(); ++c)
1148 same = (raw[c] == header[c]);
1154 while (parser.
next()) {
1155 const auto& tokens = parser.
current();
1156 if (tokens.size() != header.size())
1158 "Trajectory " << (i + 1) <<
", row " << parser.
nbLine() <<
": expected "
1159 << header.size() <<
" columns, got " << tokens.size());
1160 buffer.push_back({tokens.begin(), tokens.end()});
1163 if (buffer.size() <
k)
1165 "Trajectory " << (i + 1) <<
" has " << buffer.size()
1166 <<
" time steps but at least k=" <<
k <<
" are required");
1175 std::unordered_map< Size, std::string > atempValue;
1176 for (
const Size col: atemVarsCols)
1177 for (
const auto& r: buffer)
1178 if (!missingSet.contains(r[col])) {
1179 atempValue[col] = r[col];
1187 const auto cellAt = [&](std::size_t tt,
Size col) ->
const std::string& {
1188 if (!atemVarsCols.contains(col))
return buffer[tt][col];
1189 const auto it = atempValue.find(col);
1190 return (it == atempValue.end()) ? buffer[0][col] : it->second;
1194 for (std::size_t t = 0; t +
k <= buffer.size(); ++t) {
1198 for (
Size col = 0; col < buffer[0].size(); ++col) {
1199 row.push_back(cellAt(t, col));
1201 for (
Size slice = 1; slice <
k; ++slice) {
1202 for (
Size col = 0; col < buffer[0].size(); ++col) {
1203 if (!atemVarsCols.contains(col)) { row.push_back(buffer[t + slice][col]); }
1206 insertIfComplete(transitionTable, row);
1211 for (
Size col = 0; col < buffer[0].size(); ++col) {
1212 row.push_back(cellAt(0, col));
1214 for (
Size slice = 1; slice <
k - 1; ++slice) {
1215 for (
Size col = 0; col < buffer[0].size(); ++col) {
1216 if (!atemVarsCols.contains(col)) { row.push_back(buffer[slice][col]); }
1219 insertIfComplete(initTable, row);
1225 if (nbAtempVars > 1) {
1227 for (
Size col = 0; col < buffer[0].size(); ++col)
1228 if (atemVarsCols.contains(col)) row.push_back(cellAt(0, col));
1229 insertIfComplete(atemporalTable, row);
1239 "every row was dropped as incomplete ("
1241 <<
" in total): no fully observed transition window (or initial block) is left "
1242 "to learn from. The trajectories are too sparsely observed for k="
1253 if (nbAtempVars > 1) {
1270 const int ki = (int)
k;
1271 const auto& temporalVars =
_prior_ktbn_.temporalVarNames();
1272 const auto& atemporalVars =
_prior_ktbn_.atemporalVarNames();
1274 for (
const auto& base: temporalVars)
1275 for (
int slice = 0; slice < ki - 1; ++slice)
1278 for (
const auto& atemBase: atemporalVars) {
1288 for (
const auto& tailBase: temporalVars)
1289 for (
const auto& headBase: temporalVars)
1290 for (
int tailSlice = 1; tailSlice < ki - 1; ++tailSlice)
1291 for (
int headSlice = 0; headSlice < tailSlice; ++headSlice)
1297 template < GUM_Numeric GUM_SCALAR >
1302 template < GUM_Numeric GUM_SCALAR >
1304 const std::string b{base};
1309 template < GUM_Numeric GUM_SCALAR >
1312 const BayesNet< GUM_SCALAR >& initialBN,
1313 const BayesNet< GUM_SCALAR >& atemporalBN)
const {
1320 BayesNet< GUM_SCALAR > bn(initialBN);
1323 for (
const auto& base:
_prior_ktbn_.temporalVarNames()) {
1324 const std::string name =
_encode_(base,
k - 1);
1325 const NodeId transId = transitionBN.idFromName(name);
1326 bn.add(transitionBN.variable(transId));
1330 for (
const auto& base:
_prior_ktbn_.temporalVarNames()) {
1331 const std::string headName =
_encode_(base,
k - 1);
1332 const NodeId transHead = transitionBN.idFromName(headName);
1333 for (
const NodeId transParent: transitionBN.parents(transHead)) {
1334 const std::string& parentName = transitionBN.variable(transParent).name();
1335 bn.addArc(parentName, headName);
1344 if (atemporalBN.size() != 0) {
1345 for (
const auto& atemBase:
_prior_ktbn_.atemporalVarNames()) {
1346 const NodeId atemTail = atemporalBN.idFromName(atemBase);
1347 for (
const NodeId atemChild: atemporalBN.children(atemTail)) {
1348 const std::string& childName = atemporalBN.variable(atemChild).name();
1349 bn.addArc(atemBase, childName);
1355 for (
const auto& base:
_prior_ktbn_.temporalVarNames()) {
1356 const std::string name =
_encode_(base,
k - 1);
1357 const NodeId bnId = bn.idFromName(name);
1358 const NodeId transId = transitionBN.idFromName(name);
1359 bn.cpt(bnId).fillWith(transitionBN.cpt(transId));
1364 if (atemporalBN.size() != 0) {
1365 for (
const auto& atemBase:
_prior_ktbn_.atemporalVarNames()) {
1366 const NodeId bnId = bn.idFromName(atemBase);
1367 const NodeId atemId = atemporalBN.idFromName(atemBase);
1368 bn.cpt(bnId).fillWith(atemporalBN.cpt(atemId));
Class for fast parsing of CSV file (never more than one line in application memory).
A structure/parameter learner for k-order dynamic Bayesian networks.
void addArc(NodeId tail, NodeId head) final
insert a new arc into the directed graph
Base class for discrete random variable.
Exception: at least one argument passed to a function is not what was expected.
static constexpr int ATEMPORAL
Conventional time-slice value denoting an atemporal (static) variable.
static KTBN< GUM_SCALAR > fromBN(const BayesNet< GUM_SCALAR > &bn, const std::unordered_set< std::string > &atemporalNodes={}, std::vector< std::string > *warnings=nullptr)
Builds a k-DBN from an existing gum::BayesNet, reading its node names under one of two mutually exclu...
Error: A name of variable is not found in the database.
virtual void addNodeWithId(const NodeId id)
try to insert a node with the given id
Exception : operation not allowed.
Error: An unknown label is found in the database.
const std::string & name() const
returns the name of the variable
Class for fast parsing of CSV file (never more than one line in application memory).
bool next()
gets the next line of the csv stream and parses it
std::size_t nbLine() const
returns the current line number within the stream
const std::vector< std::string > & current() const
returns the current parsed line
the class for packing together the translators used to preprocess the datasets
DBTranslator & translatorSafe(const std::size_t k)
returns the kth translator
virtual const Variable * variable() const =0
returns the variable stored into the translator
The class representing a tabular database as used by learning tasks.
void setVariableNames(const std::vector< std::string > &names, const bool from_external_object=true) override
sets the names of the variables
std::size_t insertTranslator(const DBTranslator &translator, const std::size_t input_column, const bool unique_column=true)
insert a new translator into the database table
void reorder(const std::size_t k, const bool k_is_input_col=false)
performs a reordering of the kth translator or of the first translator parsing the kth column of the ...
void insertRow(const std::vector< std::string > &new_row) override
insert a new row at the end of the database
std::size_t nbRows() const noexcept
returns the number of records (rows) in the database
void _checkBaseIsTemporal_(std::string_view base, std::string_view context) const
Throw InvalidArgument unless base is a known temporal base. context completes "cannot appear in <cont...
void _checkArcTemporallyFeasible_(std::string_view tail, std::string_view head, std::string_view action) const
Reject an arc the k-TBN definition can never contain, so eraseForbiddenArc and addMandatoryArc both f...
static std::unordered_set< std::string > _scanConstantColumns_(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, const std::vector< std::string > &missingSymbols)
Scans every trajectory and returns the base names classified atemporal: those whose value never chang...
static void _checkMinimalOrder_(Size order, std::string_view label)
Throw InvalidArgument unless order is at least 2, label naming the offending parameter ("k" for the f...
std::string _encode_(std::string_view base, int slice) const
(base, slice) -> engine name ("A[1]" / atemporal engine name). Pure function, shared by every learner...
std::pair< std::string, int > _determineNode_(const std::string &name) const
engine name -> (base, slice); atemporal names map to KTBN::ATEMPORAL. Shared by every learner; only t...
Learns a k-TBN (structure and/or parameters) from trajectory CSVs.
static KTBN< GUM_SCALAR > _buildPriorFromCSV_(std::string_view dirPath, std::string_view csvBaseName, Size k, const std::unordered_set< std::string > &atemporalVars, const std::vector< std::string > &missingSymbols, bool induceTypes)
called in the member-initialiser list of the k-CSV constructor: opens the first trajectory CSV,...
std::vector< std::string > names() const
Base names (no slice suffix), one entry per base variable (temporal or atemporal),...
void copyState(const KTBNLearner< GUM_SCALAR > &learner)
Copy all score/algorithm/prior/constraint settings from another KTBNLearner (does not copy the databa...
KTBNLearner< GUM_SCALAR > & useNMLCorrection() override
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
const std::unordered_set< std::string > & _atemporalVarNames_() const override
atemporal base names for IKTBNLearner's shared encode/_determineNode_; read straight from the prior k...
std::unique_ptr< BNLearner< GUM_SCALAR > > _initialLearner_
learns the initial slices 0..k-2
KTBNLearner< GUM_SCALAR > & eraseForbiddenArcAllSlices(std::string_view tailBase, std::string_view headBase) override
Undo a previous addForbiddenArcAllSlices.
void _build_(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, const std::vector< std::string > &missingSymbols)
reads every trajectory, builds the three DatabaseTables (sliding window, initial-slice flattening,...
KTBNLearner< GUM_SCALAR > & erasePossibleEdge(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice) override
Undo a previous addPossibleEdge.
KTBN< GUM_SCALAR > learnKTBN() override
Full learning (structure + CPTs). Mirrors BNLearner::learnBN().
bool _ignoreMissingSymbols_
prior k-TBN: the single source of truth for k, variable domains, temporal/atemporal classification an...
std::vector< std::pair< std::string, std::string > > latentVariables() const
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
KTBNLearner< GUM_SCALAR > & addForbiddenArc(std::string_view tailNode, std::string_view headNode) override
Forbid tailNode from ever parenting headNode (engine names, e.g. "X[1]", "C").
Size nbDroppedRows() const
Number of rows dropped from the internal databases because they carried a missing symbol.
KTBNLearner< GUM_SCALAR > & useScoreAIC() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & useGreedyHillClimbing() override
static KTBN< GUM_SCALAR > _buildPriorFromBN_(Size k, const BayesNet< GUM_SCALAR > &bn, const std::unordered_set< std::string > &atemporalVars)
called in the member-initialiser list of the BN constructor: builds and returns a KTBN whose variable...
Size nbSamples() const
Number of trajectory CSV files loaded (the constructor's nbSamples).
KTBNLearner< GUM_SCALAR > & useScoreLog2Likelihood() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
std::vector< Size > nbRows() const
Number of time steps in each trajectory CSV (one entry per sample, in load order)....
bool isConstraintBased() const
True if the current structure-learning algorithm is constraint-based (e.g. MIIC).
void _forOwningLearner_(std::string_view tail, std::string_view head, F &&f)
Apply f to the ONE internal learner that can learn the arc tail -> head, chosen by its head: a slice-...
KTBN< GUM_SCALAR > _assemble_(const BayesNet< GUM_SCALAR > &transitionBN, const BayesNet< GUM_SCALAR > &initialBN, const BayesNet< GUM_SCALAR > &atemporalBN) const
glues the three parameter-learned BNs into a single k-TBN
Size nbCols() const
Number of columns in each CSV, i.e. of base variables (temporal + atemporal).
void _forEachLearner_(F &&f)
Apply f to each present internal learner (the atemporal one only when it exists). Factors out the fan...
Size _nbTemporalPossibleEdges_
counts of currently-active possible edges, split by kind: edges with at least one temporal endpoint,...
KTBNLearner< GUM_SCALAR > & addForbiddenArcAllSlices(std::string_view tailBase, std::string_view headBase) override
Forbid tailBase -> headBase at every causally-possible slice pair (every lag, not just matching slice...
std::unique_ptr< BNLearner< GUM_SCALAR > > _transitionLearner_
learns the transition kernel (arcs arriving at slice k-1)
KTBNLearner< GUM_SCALAR > & useScoreBD() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & allowArcReversals(bool allow=true) override
Allow or forbid arc reversals during structure search.
KTBNLearner< GUM_SCALAR > & addNoChildrenNode(std::string_view base, int slice) override
Declare a single (base, slice) node as a leaf (no children).
std::vector< std::size_t > domainSizes() const
Domain sizes of the base variables, in the same column order as names().
KTBN< GUM_SCALAR > learnParameters(const KTBN< GUM_SCALAR > &structure, bool takeIntoAccountScore=true)
CPTs only, using the arc structure of structure. structure must have the same base variables (names a...
Size _nbDroppedRows_
number of time steps (rows) in each trajectory CSV, in load order. rows build() dropped because they ...
KTBN< GUM_SCALAR > _prior_ktbn_
KTBNLearner< GUM_SCALAR > & allowArcDeletions(bool allow=true) override
Allow or forbid arc deletions during structure search.
KTBNLearner< GUM_SCALAR > & useExtendedGreedyHillClimbing() override
std::string checkScorePriorCompatibility() const
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & useNoCorrection() override
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
bool isScoreBased() const
True if the current structure-learning algorithm is score-based (e.g. BIC, AIC).
KTBNLearner< GUM_SCALAR > & allowArcAdditions(bool allow=true) override
Allow or forbid arc additions during structure search.
Size domainSize(std::string_view base) const
Domain size of the base variable base (e.g. "X", "C"). Engine names (e.g. "X[1]") are also accepted.
KTBNLearner< GUM_SCALAR > & eraseForbiddenArc(std::string_view tailNode, std::string_view headNode) override
Undo a previous addForbiddenArc (engine names, e.g. "X[1]", "C").
KTBNLearner< GUM_SCALAR > & useScoreBIC() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
bool isIgnoringMissingSymbols() const
Whether build() drops the rows carrying a missing symbol.
KTBNLearner< GUM_SCALAR > & useMIIC() override
void useScorefNML() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & useScoreMDL() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
KTBNLearner< GUM_SCALAR > & eraseNoChildrenNode(std::string_view base, int slice) override
Undo a previous addNoChildrenNode for a single (base, slice) node.
static std::unordered_set< std::string > _inferAtemporalVars_(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, Size k, const std::vector< std::string > &missingSymbols)
checks k >= 2, then delegates the actual scan to the shared IKTBNLearner::scanConstantColumns() (also...
KTBNLearner< GUM_SCALAR > & eraseForbiddenIntraSliceArc(std::string_view tailBase, std::string_view headBase) override
Undo a previous addForbiddenIntraSliceArc.
std::string toString() const
Human-readable summary of the learner's current configuration.
std::vector< Size > _nbTimeSlices_
Captured once by build() and exposed by nbRows(). This is the raw trajectory length,...
std::vector< std::tuple< std::string, std::string, std::string > > state() const
Settings as a vector of (key, value, comment) tuples (mirrors BNLearner::state()).
KTBNLearner< GUM_SCALAR > & addMandatoryArc(std::string_view tailNode, std::string_view headNode) override
Force tailNode to be a parent of headNode (engine names, e.g. "X[1]", "C").
bool hasMissingValues() const
True if any internal database contains missing values.
KTBNLearner< GUM_SCALAR > & setMaxIndegree(Size max_indegree) override
Cap the number of parents of any single node.
KTBNLearner< GUM_SCALAR > & addForbiddenIntraSliceArc(std::string_view tailBase, std::string_view headBase) override
Forbid tailBase -> headBase at every intra-slice position (i.e. tailBase[t] -> headBase[t] for all t ...
KTBNLearner< GUM_SCALAR > & addPossibleEdge(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice) override
Add a candidate edge for MIIC (only edges explicitly listed are explored).
bool _isKnownBase_(std::string_view base) const override
whether base is one of this learner's variables; read straight from the prior k-TBN,...
KTBNLearner< GUM_SCALAR > & useMDLCorrection() override
Engine-name pairs (tail, head) of arcs flagged as hiding a latent variable by MIIC (merged from all i...
KTBNLearner< GUM_SCALAR > & useScoreBDeu() override
Returns a warning string if the current score and prior are incompatible, empty string otherwise.
Size _nbAtemporalPossibleEdges_
KTBNLearner< GUM_SCALAR > & useSmoothingPrior(double weight=1.0) override
KTBNLearner< GUM_SCALAR > & eraseMandatoryArc(std::string_view tailNode, std::string_view headNode) override
Undo a previous addMandatoryArc (engine names, e.g. "X[1]", "C").
KTBNLearner< GUM_SCALAR > & addNoParentNode(std::string_view base, int slice) override
Declare a single (base, slice) node as a root (no parents).
KTBNLearner< GUM_SCALAR > & useLocalSearchWithTabuList(Size tabu_size=100, Size nb_decrease=2) override
std::unique_ptr< BNLearner< GUM_SCALAR > > _atemporalLearner_
learns the atemporal variables (arcs atemporal -> atemporal)
KTBNLearner(std::string_view dirPath, std::string_view csvBaseName, Size nbSamples, Size k, const std::unordered_set< std::string > &atemporalVars, const std::vector< std::string > &missingSymbols={"?"}, bool induceTypes=true, bool ignoreMissingSymbols=false)
Structure-learning constructor — variable roles supplied explicitly.
Size k() const
Order of the k-TBN being learned.
KTBNLearner< GUM_SCALAR > & eraseNoParentNode(std::string_view base, int slice) override
Undo a previous addNoParentNode for a single (base, slice) node.
void _forEachAllSlicesPair_(std::string_view tailBase, std::string_view headBase, F &&f) const
Apply f(tailSlice, headSlice) to every causally-possible slice pair of an all-slices constraint betwe...
The class representing a tabular database stored in RAM.
#define GUM_ERROR(type, msg)
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Size NodeId
Type for node ids.
include the inlined functions if necessary