70 template < std::
integral T >
77 template < GUM_Numeric GUM_SCALAR >
79 if (slice ==
ATEMPORAL)
return std::string{base};
80 return std::string{base} +
'[' + std::to_string(slice) +
']';
83 template < GUM_Numeric GUM_SCALAR >
85 const std::size_t bracketPos = name.rfind(
'[');
86 if (bracketPos == std::string_view::npos)
return {std::string{name},
ATEMPORAL};
88 const std::string_view bracketContent = name.substr(bracketPos + 1);
89 if (bracketContent.empty() || bracketContent.back() !=
']')
92 const std::string_view digits = bracketContent.substr(0, bracketContent.size() - 1);
93 if (digits.empty())
return {std::string{name},
ATEMPORAL};
94 for (
const char c: digits)
95 if (std::isdigit(
static_cast< unsigned char >(c)) == 0)
return {std::string{name},
ATEMPORAL};
99 slice = std::stoi(std::string{digits});
100 }
catch (
const std::out_of_range&) {
102 "Node name '" << name <<
"' has a slice index too large to represent as int.")
104 return {std::string{name.substr(0, bracketPos)}, slice};
107 template < GUM_Numeric GUM_SCALAR >
108 INLINE std::pair< std::string, int >
116 template < GUM_Numeric GUM_SCALAR >
117 std::vector< std::pair< std::string, int > >
119 std::vector< std::pair< std::string, int > > result;
120 result.reserve(ids.
size());
121 for (
const NodeId id: ids)
130 template < GUM_Numeric GUM_SCALAR >
133 GUM_CONSTRUCTOR(
KTBN)
136 template < GUM_Numeric GUM_SCALAR >
141 template < GUM_Numeric GUM_SCALAR >
148 template < GUM_Numeric GUM_SCALAR >
150 _k_(source._k_),
_bn_(std::move(source._bn_)),
_temporal_(std::move(source._temporal_)),
155 template < GUM_Numeric GUM_SCALAR >
157 if (
this != &source) {
167 template < GUM_Numeric GUM_SCALAR >
169 if (
this != &source) {
172 _bn_ = std::move(source._bn_);
183 template < GUM_Numeric GUM_SCALAR >
188 template < GUM_Numeric GUM_SCALAR >
193 template < GUM_Numeric GUM_SCALAR >
195 return _bn_.sizeArcs();
198 template < GUM_Numeric GUM_SCALAR >
203 template < GUM_Numeric GUM_SCALAR >
214 template < GUM_Numeric GUM_SCALAR >
219 if (ownSet.contains(base))
221 (temporal ?
"A temporal process '" :
"An atemporal variable '")
222 << base <<
"' already exists.")
223 if (otherSet.contains(base))
225 (temporal ?
"Cannot add temporal process '" :
"Cannot add atemporal variable '")
226 << base <<
"': " << (temporal ?
"an atemporal variable" :
"a temporal process")
227 <<
" with that name already exists.")
230 for (
Size t = 0; t <
_k_; ++t) {
231 const std::string encoded =
_encode_(base,
static_cast< int >(t));
234 "Temporal process '" << base <<
"' at slice " << t <<
" would produce node '"
235 << encoded <<
"' which conflicts with atemporal variable '"
239 const auto [decodedBase, decodedSlice] =
_decodeName_(base);
249 "Atemporal variable name '" << base <<
"' conflicts with temporal process '"
251 <<
"': that name is already used by its slice "
254 "Atemporal variable name '"
255 << base <<
"' is invalid: '" << decodedBase
256 <<
"' is a temporal process, so every bracket-suffixed name over it is "
257 "reserved -- including slice "
258 << decodedSlice <<
", beyond the current order k=" <<
_k_ <<
".")
263 template < GUM_Numeric GUM_SCALAR >
265 const std::string base = var.
name();
269 for (
Size t = 0; t <
_k_; ++t) {
271 std::unique_ptr< DiscreteVariable > clone(var.
clone());
272 clone->setName(
_encode_(base,
static_cast< int >(t)));
283 template < GUM_Numeric GUM_SCALAR >
286 unsigned int default_nbrmod) {
291 template < GUM_Numeric GUM_SCALAR >
296 template < GUM_Numeric GUM_SCALAR >
301 template < GUM_Numeric GUM_SCALAR >
303 unsigned int default_nbrmod) {
304 add(fast_description,
true, default_nbrmod);
307 template < GUM_Numeric GUM_SCALAR >
309 unsigned int default_nbrmod) {
310 add(fast_description,
false, default_nbrmod);
317 template < GUM_Numeric GUM_SCALAR >
319 const std::string baseStr{base};
323 template < GUM_Numeric GUM_SCALAR >
328 template < GUM_Numeric GUM_SCALAR >
333 template < GUM_Numeric GUM_SCALAR >
338 template < GUM_Numeric GUM_SCALAR >
343 template < GUM_Numeric GUM_SCALAR >
345 std::vector< std::pair< std::string, int > > result;
346 result.reserve(
size());
351 result.emplace_back(p,
static_cast< int >(t));
355 template < GUM_Numeric GUM_SCALAR >
361 template < GUM_Numeric GUM_SCALAR >
362 std::vector< std::pair< std::string, int > >
368 template < GUM_Numeric GUM_SCALAR >
374 template < GUM_Numeric GUM_SCALAR >
375 std::vector< std::pair< std::string, int > >
381 template < GUM_Numeric GUM_SCALAR >
383 const std::string baseStr{base};
397 template < GUM_Numeric GUM_SCALAR >
399 const std::string oldStr{oldBase};
400 const std::string newStr{newBase};
402 if (oldStr == newStr)
return;
414 "Renaming to '" << newStr <<
"': slice " << t
415 <<
" collides with an atemporal variable.")
417 _bn_.changeVariableName(
_encode_(oldStr,
static_cast< int >(t)),
418 _encode_(newStr,
static_cast< int >(t)));
422 const auto [decodedBase, decodedSlice] =
_decodeName_(newStr);
425 "'" << newStr <<
"' conflicts with temporal process '" << decodedBase
426 <<
"': that name is already used by its slice nodes.")
427 _bn_.changeVariableName(oldStr, newStr);
433 template < GUM_Numeric GUM_SCALAR >
435 const std::string baseStr{base};
441 "'" << baseStr <<
"' is a temporal process but is used as atemporal.")
444 return _bn_.idFromName(baseStr);
450 "'" << baseStr <<
"' is an atemporal variable but is used at slice " << slice
454 if (slice < 0 ||
Size(slice) >=
_k_)
456 "Slice " << slice <<
" is out of [0," << (
_k_ - 1) <<
"] for process '" << baseStr
461 template < GUM_Numeric GUM_SCALAR >
467 template < GUM_Numeric GUM_SCALAR >
473 template < GUM_Numeric GUM_SCALAR >
479 template < GUM_Numeric GUM_SCALAR >
489 template < GUM_Numeric GUM_SCALAR >
492 std::string_view headBase,
500 "A temporal variable cannot be a parent of the atemporal variable '" << headBase
502 }
else if (tailSlice !=
ATEMPORAL && tailSlice > headSlice) {
504 "An arc cannot go from a future slice (" << tailSlice <<
") to a past slice ("
505 << headSlice <<
").")
508 _bn_.addArc(tail, head);
511 template < GUM_Numeric GUM_SCALAR >
514 std::string_view headBase,
519 template < GUM_Numeric GUM_SCALAR >
522 std::string_view headBase,
523 int headSlice)
const {
528 template < GUM_Numeric GUM_SCALAR >
535 template < GUM_Numeric GUM_SCALAR >
542 template < GUM_Numeric GUM_SCALAR >
549 template < GUM_Numeric GUM_SCALAR >
550 std::vector< std::pair< std::pair< std::string, int >, std::pair< std::string, int > > >
552 std::vector< std::pair< std::pair< std::string, int >, std::pair< std::string, int > > > result;
553 result.reserve(
_bn_.sizeArcs());
554 for (
const auto& arc:
_bn_.arcs()) {
565 template < GUM_Numeric GUM_SCALAR >
571 template < GUM_Numeric GUM_SCALAR >
577 template < GUM_Numeric GUM_SCALAR >
582 template < GUM_Numeric GUM_SCALAR >
587 template < GUM_Numeric GUM_SCALAR >
593 template < GUM_Numeric GUM_SCALAR >
595 std::string_view base,
598 const std::vector< GUM_SCALAR >& distribution)
const {
600 const Tensor< GUM_SCALAR >&
cpt =
_bn_.cpt(
id);
605 "fillCPT: distribution has " << distribution.size() <<
" value(s) but '" << base
606 <<
"' has " << self.
domainSize() <<
" modalities.")
608 const Size nbParents =
cpt.nbrDim() - 1;
609 if (
parents.size() != nbParents)
611 "fillCPT: " <<
parents.size() <<
" parent value(s) given but the node has "
612 << nbParents <<
" parent(s); every parent must be specified.")
619 for (
const auto& [parNode, parVal]:
parents) {
620 const auto& [parBase, parSlice] = parNode;
623 if (parId ==
id || !
cpt.contains(parVar))
625 "fillCPT: '" << parBase <<
"' is not a parent of the target node.")
626 inst.
chgVal(parVar, parVal.toIndex(parVar));
632 cpt.set(inst, distribution[m]);
636 template < GUM_Numeric GUM_SCALAR >
638 std::string_view node_name,
639 const std::map< std::variant< std::string, std::pair< std::string, int > >,
KTBNModality >&
641 const std::vector< GUM_SCALAR >& distribution)
const {
642 const NodeId id =
_bn_.idFromName(std::string{node_name});
643 const Tensor< GUM_SCALAR >&
cpt =
_bn_.cpt(
id);
648 "fillCPT: distribution has " << distribution.size() <<
" value(s) but '"
652 const Size nbParents =
cpt.nbrDim() - 1;
653 if (
parents.size() != nbParents)
655 "fillCPT: " <<
parents.size() <<
" parent value(s) given but the node has "
656 << nbParents <<
" parent(s); every parent must be specified.")
664 for (
const auto& [parKey, parVal]:
parents) {
665 const std::string parName
666 = std::holds_alternative< std::string >(parKey)
667 ? std::get< std::string >(parKey)
668 :
_encode_(std::get< std::pair< std::string, int > >(parKey).first,
669 std::get< std::pair< std::string, int > >(parKey).second);
670 const NodeId parId =
_bn_.idFromName(parName);
672 if (parId ==
id || !
cpt.contains(parVar))
674 "fillCPT: '" << parName <<
"' is not a parent of the target node.")
678 inst.
chgVal(parVar, parVal.toIndex(parVar));
683 cpt.set(inst, distribution[m]);
691 template < GUM_Numeric GUM_SCALAR >
693 return BayesNet< GUM_SCALAR >(
_bn_);
696 template < GUM_Numeric GUM_SCALAR >
698 if (nbTimeSlices <
_k_)
700 "Cannot unroll over " << nbTimeSlices <<
" slices: fewer than the order k=" <<
_k_
703 BayesNet< GUM_SCALAR > unrolled;
704 const int kernelSlice =
static_cast< int >(
_k_ - 1);
708 unrolled.add(
_bn_.variable(
_bn_.idFromName(a)));
715 for (
Size t = 0; t < nbTimeSlices; ++t) {
717 std::unique_ptr< DiscreteVariable > clone(templateVar.
clone());
718 clone->setName(
_encode_(p,
static_cast< int >(t)));
719 unrolled.add(*clone);
725 for (
const auto& arc:
_bn_.arcs()) {
726 unrolled.addArc(
_bn_.variable(arc.tail()).name(),
_bn_.variable(arc.head()).name());
731 unrolled.cpt(
_bn_.variable(n).name()).fillWith(
_bn_.cpt(n));
738 const Tensor< GUM_SCALAR >& templateCpt =
_bn_.cpt(lastSliceNodeId);
741 std::vector< std::pair< std::string, int > > lags;
742 for (
const auto& [parBase, parSlice]:
parents(p, kernelSlice)) {
744 lags.emplace_back(parBase, lag);
748 std::vector< std::string > templateVarNames;
750 for (
Size t =
_k_; t < nbTimeSlices; ++t) {
751 const std::string child =
_encode_(p,
static_cast< int >(t));
754 unrolledToTemplate.
clear();
756 for (
const auto& [parBase, lag]: lags) {
758 unrolled.addArc(parBase, child);
759 unrolledToTemplate.
insert(parBase, parBase);
761 const std::string parName =
_encode_(parBase,
static_cast< int >(t) - lag);
762 unrolled.addArc(parName, child);
763 unrolledToTemplate.
insert(parName,
_encode_(parBase, kernelSlice - lag));
768 const Tensor< GUM_SCALAR >& unrolledCpt = unrolled.cpt(child);
769 templateVarNames.clear();
770 templateVarNames.reserve(unrolledCpt.nbrDim());
771 for (
Idx i = 0; i < unrolledCpt.nbrDim(); ++i) {
772 templateVarNames.push_back(unrolledToTemplate[unrolledCpt.variable(i).name()]);
774 unrolledCpt.fillWith(templateCpt, templateVarNames);
785 template < GUM_Numeric GUM_SCALAR >
789 std::string filepath{filename};
790 const bool text = filepath.ends_with(
".jgum");
791 if (!text && !filepath.ends_with(
".bgum")) filepath +=
".bgum";
792 return {std::move(filepath), !text};
795 template < GUM_Numeric GUM_SCALAR >
803 auto join = [](
const std::unordered_set< std::string >& names) {
805 for (
const auto& n: names) {
806 if (!out.empty()) out +=
',';
807 for (
const char c: n) {
808 if (c ==
'\\' || c ==
',') out +=
'\\';
815 BayesNet< GUM_SCALAR > annotated =
_bn_;
816 annotated.setProperty(
"KTBN.k", std::to_string(
_k_));
817 annotated.setProperty(
"KTBN.temporal", join(
_temporal_));
818 annotated.setProperty(
"KTBN.atemporal", join(
_atemporal_));
820 writer.
write(filepath, annotated);
823 template < GUM_Numeric GUM_SCALAR >
827 BayesNet< GUM_SCALAR > bn;
831 std::stringstream stream;
839 if (!(bn.existsProperty(
"KTBN.k") && bn.existsProperty(
"KTBN.temporal")
840 && bn.existsProperty(
"KTBN.atemporal")))
843 auto split = [](std::string_view csv, std::unordered_set< std::string >& out) {
845 bool escaped =
false;
846 for (
const char c: csv) {
850 }
else if (c ==
'\\') {
852 }
else if (c ==
',') {
853 if (!current.empty()) out.insert(current);
859 if (!current.empty()) out.insert(current);
864 k_val =
static_cast< Size >(std::stoul(bn.property(
"KTBN.k")));
865 }
catch (
const std::exception& e) {
867 "KTBN::load: malformed KTBN.k property ('" << bn.property(
"KTBN.k")
868 <<
"'): " << e.what())
870 KTBN< GUM_SCALAR > res(k_val);
872 split(bn.property(
"KTBN.temporal"), res._temporal_);
873 split(bn.property(
"KTBN.atemporal"), res._atemporal_);
878 template < GUM_Numeric GUM_SCALAR >
881 const std::unordered_set< std::string >& atemporalNodes,
882 std::vector< std::string >* warnings) {
883 KTBN< GUM_SCALAR > res(1);
885 res._determineNodesFromBN_(atemporalNodes, warnings);
889 template < GUM_Numeric GUM_SCALAR >
891 const std::unordered_set< std::string >& atemporalNodes,
892 std::vector< std::string >* warnings) {
898 for (
const std::string& name: atemporalNodes)
899 if (!
_bn_.exists(name))
902 const auto warn = [warnings](
const std::string& message) {
903 if (warnings !=
nullptr) warnings->push_back(message);
912 const auto decodeTrailingSlice = [](std::string_view name) -> std::pair< std::string, int > {
913 std::size_t pos = name.size();
914 while (pos > 0 && std::isdigit(
static_cast< unsigned char >(name[pos - 1])))
916 if (pos == name.size())
return {std::string{name},
ATEMPORAL};
918 const std::string_view digits = name.substr(pos);
921 slice = std::stoi(std::string{digits});
922 }
catch (
const std::out_of_range&) {
924 "Node name '" << name <<
"' has a slice index too large to represent as int.")
926 return {std::string{name.substr(0, pos)}, slice};
938 bool hasBracket =
false;
940 const std::string& name =
_bn_.variable(n).name();
941 if (atemporalNodes.contains(name))
continue;
949 std::vector< std::string > discovered;
954 const std::string& name =
_bn_.variable(n).name();
959 if (atemporalNodes.contains(name)) {
964 const auto [base, slice] = hasBracket ?
_decodeName_(name) : decodeTrailingSlice(name);
969 if (!slicesPerProcess.
exists(base)) {
971 discovered.push_back(base);
974 if (slicesPerProcess[base].
exists(slice))
976 "Two variables map to process '" << base <<
"' at slice " << slice <<
".")
978 slicesPerProcess[base].
insert(slice, n);
980 if (slice > maxSlice) maxSlice = slice;
984 _k_ = (maxSlice < 0) ?
Size(1) :
Size(maxSlice + 1);
991 const char*
const conventionNoun = hasBracket ?
"bracket" :
"digit-suffixed";
992 for (
const auto& base: discovered) {
997 "Base name '" << base <<
"' is used both as an atemporal variable (bare node '"
998 << base <<
"') and as a temporal process (via " << conventionNoun
999 <<
" nodes). " <<
"Rename one of them before calling fromBN().")
1001 std::string missing;
1002 for (
Size t = 0; t <
_k_; ++t)
1003 if (!sliceMap.
exists(
static_cast< int >(t))) {
1004 if (!missing.empty()) missing +=
", ";
1005 missing += std::to_string(t);
1008 if (missing.empty()) {
1018 std::unique_ptr< DiscreteVariable > probe(ref.
clone());
1019 for (
Size t = 1; t <
_k_; ++t) {
1021 probe->setName(other.
name());
1022 if (!(*probe == other))
1024 "The temporal process '"
1025 << base <<
"' has mismatched slice variables: '" << ref.
name() <<
"' is "
1027 <<
". Every slice of a process must have the same type and domain.")
1036 for (
Size t = 0; t <
_k_; ++t)
1037 _bn_.changeVariableName(
_bn_.variable(sliceMap[
static_cast< int >(t)]).name(),
1038 _encode_(base,
static_cast< int >(t)));
1043 std::string reclassified;
1044 for (
auto it = sliceMap.
cbegin(); it != sliceMap.
cend(); ++it) {
1045 const std::string& nodeName =
_bn_.variable(it.val()).name();
1047 if (!reclassified.empty()) reclassified +=
", ";
1048 reclassified +=
"'" + nodeName +
"'";
1050 warn(
"Node(s) " + reclassified +
" look temporal (base='" + base +
"', " + conventionNoun
1051 +
" convention) but the process is missing slice(s) " + missing
1052 +
" for k=" + std::to_string(
_k_)
1053 +
": they are classified as atemporal variables, original name kept. Pass them in "
1054 "fromBN()'s atemporalNodes argument to make that explicit and silence this warning.");
1065 for (
const auto& arc:
_bn_.arcs()) {
1070 "The network has a temporal->atemporal arc into '"
1071 <<
_bn_.variable(arc.head()).name() <<
"'.")
1074 "The network has a future->past arc " <<
_bn_.variable(arc.tail()).name() <<
"->"
1075 <<
_bn_.variable(arc.head()).name() <<
".")
1083 template < GUM_Numeric GUM_SCALAR >
1085 const auto join = [](
const std::vector< std::string >& v) {
1087 for (
const auto& n: v) {
1088 if (!s.empty()) s +=
", ";
1098 for (
const auto& arc:
_bn_.arcs())
1099 arcs += std::format(
" {} -> {}\n",
1100 _bn_.variable(arc.tail()).name(),
1101 _bn_.variable(arc.head()).name());
1103 return std::format(
"k-TBN (k={}, {} nodes, {} arcs)\n"
1104 " temporal processes ({}): {}\n"
1105 " atemporal variables ({}): {}\n"
1106 "\n arcs ({}):\n{}",
1111 join(temporalNames),
1113 join(atemporalNames),
1118 template < GUM_Numeric GUM_SCALAR >
1123 template < GUM_Numeric GUM_SCALAR >
1130 template < GUM_Numeric GUM_SCALAR >
1133 out.reserve(name.size());
1134 for (
const char c: name) {
1135 if (c ==
'"') out +=
'\\';
1145 template < GUM_Numeric GUM_SCALAR >
1147 bool highlightReplicated)
const {
1151 std::map< int, std::vector< std::pair< std::string, std::string > > > timeslices;
1152 for (
const NodeId n: bn.nodes()) {
1153 const std::string& name = bn.variable(n).name();
1155 timeslices[slice].emplace_back(name, base);
1158 std::stringstream dot;
1159 dot <<
"digraph KTBN {\n";
1160 dot <<
" rankdir=LR;\n";
1161 dot <<
" splines=ortho;\n";
1162 dot <<
" node [color=\"#000000\", fillcolor=white, style=filled];\n\n";
1164 for (
auto& [slice,
nodes]: timeslices) {
1167 dot <<
" subgraph cluster_atemporal {\n";
1168 dot <<
" label=\"atemporal\";\n";
1169 dot <<
" style=filled;\n";
1170 dot <<
" bgcolor=\"lightyellow\";\n";
1171 for (
const auto& [full, label]:
nodes)
1175 const bool replicated = highlightReplicated &&
Size(slice) >=
_k_;
1176 dot <<
" subgraph cluster_" << slice <<
" {\n";
1177 dot <<
" label=\"Time slice " << slice <<
"\";\n";
1178 dot <<
" style=filled;\n";
1179 dot <<
" bgcolor=\"" << (replicated ?
"lightcyan" :
"#DDDDDD") <<
"\";\n";
1180 for (
const auto& [full, label]:
nodes)
1187 dot <<
" edge [color=black, constraint=false];\n";
1188 for (
const auto& arc: bn.arcs())
1189 dot <<
" \"" <<
_escapeDot_(bn.variable(arc.tail()).name()) <<
"\" -> \""
1190 <<
_escapeDot_(bn.variable(arc.head()).name()) <<
"\";\n";
1192 dot <<
"\n edge [style=invis, constraint=true];\n";
1193 if (
const auto it0 = timeslices.find(0); it0 != timeslices.end()) {
1194 for (
const auto& node0: it0->second) {
1195 const std::string& label = node0.second;
1198 for (
const auto& [slice,
nodes]: timeslices) {
1213 template < GUM_Numeric GUM_SCALAR >
1215 return _bn_.toDot();
1218 template < GUM_Numeric GUM_SCALAR >
1223 const int lastSlice =
static_cast< int >(
_k_) - 1;
1224 std::set< std::tuple< std::string, std::string, int > > edges;
1225 for (
const auto& arc:
_bn_.arcs()) {
1226 const auto [tailBase, tailSlice] =
_decodeName_(
_bn_.variable(arc.tail()).name());
1227 const auto [headBase, headSlice] =
_decodeName_(
_bn_.variable(arc.head()).name());
1228 if (headSlice != lastSlice)
continue;
1231 edges.emplace(tailBase, headBase, lag);
1234 std::stringstream dot;
1235 dot <<
"digraph KTBN {\n";
1236 dot <<
" rankdir=LR;\n";
1237 dot <<
" node [color=\"#000000\", fillcolor=white, style=filled];\n\n";
1239 for (
const auto& base: baseNames)
1243 for (
const auto& [tailBase, headBase, lag]: edges) {
1245 if (lag !=
ATEMPORAL) dot <<
" [label=\"" << lag <<
"\"]";
1253 template < GUM_Numeric GUM_SCALAR >
1254 std::ostream&
operator<<(std::ostream& output,
const KTBN< GUM_SCALAR >& kdbn) {
1255 output << kdbn.toString();
Definition of classe for GUM (json) file output manipulation.
Class representing k-order dynamic Bayesian networks (k-DBN).
void write(std::ostream &output, IBayesNet< GUM_SCALAR > &bn)
Writes a Bayesian network in the output stream.
Base class for discrete random variable.
DiscreteVariable * clone() const override=0
Copy Factory.
virtual Size domainSize() const =0
std::string domain() const override=0
string represent the domain of the variable
Exception : a similar label already exists.
void showErrorCounts(std::ostream &stream=std::cerr) const
Size proceed() final
Parse the file given at construction.
void showElegantErrorsAndWarnings(std::ostream &stream=std::cerr) const
Writes a IBayesNet in the GUM json format.
const const_iterator & cend() const noexcept
Returns the unsafe const_iterator pointing to the end of the hashtable.
value_type & insert(const Key &key, const Val &val)
Adds a new element (actually a copy of this element) into the hash table.
void clear()
Removes all the elements in the hash table.
bool exists(const Key &key) const
Checks whether there exists an element with a given key in the hashtable.
const_iterator cbegin() const
Returns an unsafe const_iterator pointing to the beginning of the hashtable.
Exception : input/output problem.
Class for assigning/browsing values to tuples of discrete variables.
Instantiation & chgVal(const DiscreteVariable &v, Idx newval)
Assign newval to variable v in the Instantiation.
Exception: at least one argument passed to a function is not what was expected.
std::vector< std::pair< std::string, int > > children(std::string_view base, int slice) const
Children of a node as (base, slice) pairs (ATEMPORAL if atemporal).
void clear()
Removes all variables and arcs, keeping the order .
void addTemporal(const DiscreteVariable &var)
Convenience shortcut for add(var, true).
const std::unordered_set< std::string > & temporalVarNames() const
void addAtemporal(const DiscreteVariable &var)
Convenience shortcut for add(var, false).
BayesNet< GUM_SCALAR > toBN() const
void addArc(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice)
Adds an arc between two (process, slice) endpoints.
std::vector< std::pair< std::pair< std::string, int >, std::pair< std::string, int > > > arcs() const
std::pair< std::string, int > _decodeName_(std::string_view name) const
Purely syntactic parse of an engine name → (base, slice). Slice is ATEMPORAL when there is no [digits...
std::vector< std::pair< std::string, int > > nodes() const
BayesNet< GUM_SCALAR > _bn_
The underlying Bayesian network used as a storage engine for the template.
static KTBN< GUM_SCALAR > load(std::string_view filename)
Loads a k-DBN from a GUM file produced by save().
void erase(std::string_view base)
Removes a variable and all its incident arcs.
const DiscreteVariable & variable(std::string_view base, int slice) const
Returns the gum::DiscreteVariable of a (process, slice) couple.
std::string summaryGraph() const
Returns the Graphviz DOT string of the summary graph: the projection of the transition kernel alone (...
bool existsArc(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice) const
void save(std::string_view filename) const
Saves the template in the GUM format (text .jgum or binary .bgum).
static std::string _escapeDot_(std::string_view name)
Escapes double quotes for a DOT identifier or label. Shared by timeSlicesToDot() and summaryGraph().
NodeId _validateVariable_(std::string_view base, int slice) const
Resolves and validates a (base, slice) endpoint into its NodeId.
std::unordered_set< std::string > _temporal_
Base names of the registered temporal processes.
void eraseArc(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice)
Removes an arc between two (process, slice) endpoints.
std::vector< std::pair< std::string, int > > parents(std::string_view base, int slice) const
Parents of a node as (base, slice) pairs (ATEMPORAL if atemporal).
const std::unordered_set< std::string > & atemporalVarNames() const
void generateCPT(std::string_view base, int slice) const
Randomly generates the CPT of a single node.
Size nbAtemporalVars() const
std::string _encode_(std::string_view base, int slice) const
Encodes (base, slice) → engine name: base[t], or base if atemporal.
const Tensor< GUM_SCALAR > & cpt(std::string_view base, int slice) const
Returns the CPT of a (process, slice) couple.
void generateCPTs() const
Randomly generates every CPT of the template.
std::unordered_set< std::string > _atemporal_
Base names of the registered atemporal variables.
std::string baseName(const DiscreteVariable &var) const
Returns the base name (without bracket encoding) of var.
static constexpr int ATEMPORAL
Conventional time-slice value denoting an atemporal (static) variable.
std::string toUnrolledDot(Size T, bool highlightReplicated=false) const
Returns a Graphviz DOT string of the k-DBN unrolled over T time slices.
std::string bnToDot() const
Returns the Graphviz DOT string of the underlying storage BayesNet.
KTBN< GUM_SCALAR > & operator=(const KTBN< GUM_SCALAR > &source)
Copy assignment operator.
int timeSlice(const DiscreteVariable &var) const
The time slice of var, or ATEMPORAL if it is atemporal.
void changeVariableName(std::string_view oldBase, std::string_view newBase)
Renames a variable (temporal process or atemporal variable).
std::string toDot() const
Returns a Graphviz DOT string with one cluster per time slice.
virtual ~KTBN()
Destructor.
std::pair< std::string, int > _determineNode_(const std::string &name) const
Cache-aware classification of a node name → (base, slice): nodes registered in _atemporal_ (atemporal...
BayesNet< GUM_SCALAR > unroll(Size nbTimeSlices) const
Unrolls the k-DBN into a standard gum::BayesNet.
KTBN(Size k=2)
Default constructor.
void fillCPT(std::string_view base, int slice, const std::map< std::pair< std::string, int >, KTBNModality > &parents, const std::vector< GUM_SCALAR > &distribution) const
Fills one conditional distribution P(node | parent configuration).
std::string _timeSlicesToDot_(const BayesNet< GUM_SCALAR > &bn, bool highlightReplicated) const
Renders bn as time-slice-clustered DOT. Shared engine behind toDot() (on _bn_) and toUnrolledDot() (o...
Size _k_
The order (number of time slices in the template).
Size nbTemporalVars() const
std::vector< std::pair< std::string, int > > _determineNodeSet_(const NodeSet &ids) const
Maps a set of node ids to (base, slice) pairs (via determineNode).
bool exists(std::string_view base) const
void _validateAdd_(const std::string &base, bool temporal) const
Checks that a variable named base can be added.
std::string toString() const
void add(const DiscreteVariable &var, bool temporal=true)
Adds a variable to the k-DBN.
static std::pair< std::string, bool > _resolveGumFormat_(std::string_view filename)
Resolves a user filename to (filepath, binary): ensures a .jgum/.bgum extension (....
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...
void _determineNodesFromBN_(const std::unordered_set< std::string > &atemporalNodes, std::vector< std::string > *warnings)
Rebuilds the cached name sets from the storage engine content (used by fromBN()/load(); decodes names...
Exception : the element we looked for cannot be found.
Exception : operation not allowed.
Exception : out of bound.
bool contains(const Key &k) const
Indicates whether a given elements belong to the set.
void insert(const Key &k)
Inserts a new element into the set.
Size size() const noexcept
Returns the number of elements in the set.
Exception : problem with size.
const std::string & name() const
returns the name of the variable
#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.
Set< NodeId > NodeSet
Some typdefs and define for shortcuts ...
std::vector< std::string > split(std::string_view str, std::string_view delim)
Split str using the delimiter.
gum is the global namespace for all aGrUM entities
std::ostream & operator<<(std::ostream &stream, const AVLTree< Val, Cmp > &tree)
display the content of a tree
std::unique_ptr< DiscreteVariable > fastVariable(std::string var_description, Size default_domain_size)
Create a pointer on a Discrete Variable from a "fast" syntax.
A parent's value in gum::KTBN::fillCPT(): a modality index or a modality label.
Idx index
The index, when isLabel is false.
bool isLabel
Whether the value was spelled as a label rather than an index.
KTBNModality(T modality)
From a modality index.