aGrUM 3.2.0
a C++ library for (probabilistic) graphical models
KTBN_tpl.h
Go to the documentation of this file.
1/****************************************************************************
2 * This file is part of the aGrUM/pyAgrum library. *
3 * *
4 * Copyright (c) 2005-2026 by *
5 * - Pierre-Henri WUILLEMIN(_at_LIP6) *
6 * - Christophe GONZALES(_at_AMU) *
7 * *
8 * The aGrUM/pyAgrum library is free software; you can redistribute it *
9 * and/or modify it under the terms of either : *
10 * *
11 * - the GNU Lesser General Public License as published by *
12 * the Free Software Foundation, either version 3 of the License, *
13 * or (at your option) any later version, *
14 * - the MIT license (MIT), *
15 * - or both in dual license, as here. *
16 * *
17 * (see https://agrum.gitlab.io/articles/dual-licenses-lgplv3mit.html) *
18 * *
19 * This aGrUM/pyAgrum library is distributed in the hope that it will be *
20 * useful, but WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, *
21 * INCLUDING BUT NOT LIMITED TO THE WARRANTIES MERCHANTABILITY or FITNESS *
22 * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE *
23 * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER *
24 * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, *
25 * ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR *
26 * OTHER DEALINGS IN THE SOFTWARE. *
27 * *
28 * See LICENCES for more details. *
29 * *
30 * SPDX-FileCopyrightText: Copyright 2005-2026 *
31 * - Pierre-Henri WUILLEMIN(_at_LIP6) *
32 * - Christophe GONZALES(_at_AMU) *
33 * SPDX-License-Identifier: LGPL-3.0-or-later OR MIT *
34 * *
35 * Contact : info_at_agrum_dot_org *
36 * homepage : http://agrum.gitlab.io *
37 * gitlab : https://gitlab.com/agrumery/agrum *
38 * *
39 ****************************************************************************/
40
41
48
49#pragma once
50
51#include <algorithm>
52#include <cctype>
53#include <format>
54#include <map>
55#include <set>
56#include <sstream>
57#include <tuple>
58
62#include <agrum/KTBN/KTBN.h>
63
64namespace gum {
65
66 // ===========================================================================
67 // KTBNModality
68 // ===========================================================================
69
70 template < std::integral T >
71 KTBNModality::KTBNModality(T modality) : isLabel(false), index(static_cast< Idx >(modality)) {}
72
73 // ===========================================================================
74 // Private naming helpers (the only place names are produced / parsed)
75 // ===========================================================================
76
77 template < GUM_Numeric GUM_SCALAR >
78 INLINE std::string KTBN< GUM_SCALAR >::_encode_(std::string_view base, int slice) const {
79 if (slice == ATEMPORAL) return std::string{base};
80 return std::string{base} + '[' + std::to_string(slice) + ']';
81 }
82
83 template < GUM_Numeric GUM_SCALAR >
84 std::pair< std::string, int > KTBN< GUM_SCALAR >::_decodeName_(std::string_view name) const {
85 const std::size_t bracketPos = name.rfind('[');
86 if (bracketPos == std::string_view::npos) return {std::string{name}, ATEMPORAL};
87
88 const std::string_view bracketContent = name.substr(bracketPos + 1);
89 if (bracketContent.empty() || bracketContent.back() != ']')
90 return {std::string{name}, ATEMPORAL};
91
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};
96
97 int slice{};
98 try {
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.")
103 }
104 return {std::string{name.substr(0, bracketPos)}, slice};
105 }
106
107 template < GUM_Numeric GUM_SCALAR >
108 INLINE std::pair< std::string, int >
109 KTBN< GUM_SCALAR >::_determineNode_(const std::string& name) const {
110 // Atemporal and orphan-bracket nodes are registered in _atemporal_ and map to ATEMPORAL;
111 // any other name is a temporal "base[t]" parsed by _decodeName_.
112 if (_atemporal_.contains(name)) return {name, ATEMPORAL};
113 return _decodeName_(name);
114 }
115
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)
122 result.push_back(_determineNode_(_bn_.variable(id).name()));
123 return result;
124 }
125
126 // ===========================================================================
127 // Constructors and destructor
128 // ===========================================================================
129
130 template < GUM_Numeric GUM_SCALAR >
132 if (k == 0) GUM_ERROR(InvalidArgument, "A k-DBN must have an order k >= 1.")
133 GUM_CONSTRUCTOR(KTBN)
134 }
135
136 template < GUM_Numeric GUM_SCALAR >
138 GUM_DESTRUCTOR(KTBN)
139 }
140
141 template < GUM_Numeric GUM_SCALAR >
142 KTBN< GUM_SCALAR >::KTBN(const KTBN< GUM_SCALAR >& source) :
143 _k_(source._k_), _bn_(source._bn_), _temporal_(source._temporal_),
144 _atemporal_(source._atemporal_) {
145 GUM_CONS_CPY(KTBN)
146 }
147
148 template < GUM_Numeric GUM_SCALAR >
149 KTBN< GUM_SCALAR >::KTBN(KTBN< GUM_SCALAR >&& source) noexcept :
150 _k_(source._k_), _bn_(std::move(source._bn_)), _temporal_(std::move(source._temporal_)),
151 _atemporal_(std::move(source._atemporal_)) {
152 GUM_CONS_MOV(KTBN)
153 }
154
155 template < GUM_Numeric GUM_SCALAR >
156 KTBN< GUM_SCALAR >& KTBN< GUM_SCALAR >::operator=(const KTBN< GUM_SCALAR >& source) {
157 if (this != &source) {
158 GUM_OP_CPY(KTBN);
159 _k_ = source._k_;
160 _bn_ = source._bn_;
161 _temporal_ = source._temporal_;
162 _atemporal_ = source._atemporal_;
163 }
164 return *this;
165 }
166
167 template < GUM_Numeric GUM_SCALAR >
168 KTBN< GUM_SCALAR >& KTBN< GUM_SCALAR >::operator=(KTBN< GUM_SCALAR >&& source) noexcept {
169 if (this != &source) {
170 GUM_OP_MOV(KTBN);
171 _k_ = source._k_;
172 _bn_ = std::move(source._bn_);
173 _temporal_ = std::move(source._temporal_);
174 _atemporal_ = std::move(source._atemporal_);
175 }
176 return *this;
177 }
178
179 // ===========================================================================
180 // Accessors
181 // ===========================================================================
182
183 template < GUM_Numeric GUM_SCALAR >
185 return _k_;
186 }
187
188 template < GUM_Numeric GUM_SCALAR >
190 return _bn_.size();
191 }
192
193 template < GUM_Numeric GUM_SCALAR >
195 return _bn_.sizeArcs();
196 }
197
198 template < GUM_Numeric GUM_SCALAR >
199 INLINE bool KTBN< GUM_SCALAR >::empty() const {
200 return _bn_.empty();
201 }
202
203 template < GUM_Numeric GUM_SCALAR >
205 _bn_.clear();
206 _temporal_.clear();
207 _atemporal_.clear();
208 }
209
210 // ===========================================================================
211 // Variable management
212 // ===========================================================================
213
214 template < GUM_Numeric GUM_SCALAR >
215 void KTBN< GUM_SCALAR >::_validateAdd_(const std::string& base, bool temporal) const {
216 const auto& ownSet = temporal ? _temporal_ : _atemporal_;
217 const auto& otherSet = temporal ? _atemporal_ : _temporal_;
218
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.")
228
229 if (temporal) {
230 for (Size t = 0; t < _k_; ++t) {
231 const std::string encoded = _encode_(base, static_cast< int >(t));
232 if (_atemporal_.contains(encoded))
234 "Temporal process '" << base << "' at slice " << t << " would produce node '"
235 << encoded << "' which conflicts with atemporal variable '"
236 << encoded << "'.")
237 }
238 } else {
239 const auto [decodedBase, decodedSlice] = _decodeName_(base);
240 if (decodedSlice != ATEMPORAL && _temporal_.contains(decodedBase)) {
241 // Bracket notation over an existing temporal process is reserved in
242 // FULL, whatever the index: _decodeName_ has no upper bound, so
243 // "X[999]" decodes to (X, 999) and would shadow the process even though
244 // no such node exists. Only the in-range case can claim an actual node
245 // collision -- promising one for an out-of-range index would send the
246 // caller looking for a node that was never there.
247 if (Size(decodedSlice) < _k_)
249 "Atemporal variable name '" << base << "' conflicts with temporal process '"
250 << decodedBase
251 << "': that name is already used by its slice "
252 "nodes.")
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_ << ".")
259 }
260 }
261 }
262
263 template < GUM_Numeric GUM_SCALAR >
264 void KTBN< GUM_SCALAR >::add(const DiscreteVariable& var, bool temporal) {
265 const std::string base = var.name();
266
267 if (temporal) {
268 _validateAdd_(base, true);
269 for (Size t = 0; t < _k_; ++t) {
270 // clone() to rename; BayesNet::add() clones again internally (unavoidable via public API).
271 std::unique_ptr< DiscreteVariable > clone(var.clone());
272 clone->setName(_encode_(base, static_cast< int >(t)));
273 _bn_.add(*clone);
274 }
275 _temporal_.insert(base);
276 } else {
277 _validateAdd_(base, false);
278 _bn_.add(var);
279 _atemporal_.insert(base);
280 }
281 }
282
283 template < GUM_Numeric GUM_SCALAR >
284 void KTBN< GUM_SCALAR >::add(std::string_view fast_description,
285 bool temporal,
286 unsigned int default_nbrmod) {
287 auto v = fastVariable< GUM_SCALAR >(std::string{fast_description}, Size(default_nbrmod));
288 add(*v, temporal);
289 }
290
291 template < GUM_Numeric GUM_SCALAR >
293 add(var, true);
294 }
295
296 template < GUM_Numeric GUM_SCALAR >
298 add(var, false);
299 }
300
301 template < GUM_Numeric GUM_SCALAR >
302 INLINE void KTBN< GUM_SCALAR >::addTemporal(std::string_view fast_description,
303 unsigned int default_nbrmod) {
304 add(fast_description, true, default_nbrmod);
305 }
306
307 template < GUM_Numeric GUM_SCALAR >
308 INLINE void KTBN< GUM_SCALAR >::addAtemporal(std::string_view fast_description,
309 unsigned int default_nbrmod) {
310 add(fast_description, false, default_nbrmod);
311 }
312
313 // ===========================================================================
314 // Variable queries
315 // ===========================================================================
316
317 template < GUM_Numeric GUM_SCALAR >
318 INLINE bool KTBN< GUM_SCALAR >::exists(std::string_view base) const {
319 const std::string baseStr{base};
320 return _temporal_.contains(baseStr) || _atemporal_.contains(baseStr);
321 }
322
323 template < GUM_Numeric GUM_SCALAR >
324 INLINE const std::unordered_set< std::string >& KTBN< GUM_SCALAR >::temporalVarNames() const {
325 return _temporal_;
326 }
327
328 template < GUM_Numeric GUM_SCALAR >
329 INLINE const std::unordered_set< std::string >& KTBN< GUM_SCALAR >::atemporalVarNames() const {
330 return _atemporal_;
331 }
332
333 template < GUM_Numeric GUM_SCALAR >
335 return _temporal_.size();
336 }
337
338 template < GUM_Numeric GUM_SCALAR >
340 return _atemporal_.size();
341 }
342
343 template < GUM_Numeric GUM_SCALAR >
344 std::vector< std::pair< std::string, int > > KTBN< GUM_SCALAR >::nodes() const {
345 std::vector< std::pair< std::string, int > > result;
346 result.reserve(size());
347 for (const auto& a: _atemporal_)
348 result.emplace_back(a, ATEMPORAL);
349 for (const auto& p: _temporal_)
350 for (Size t = 0; t < _k_; ++t)
351 result.emplace_back(p, static_cast< int >(t));
352 return result;
353 }
354
355 template < GUM_Numeric GUM_SCALAR >
356 std::vector< std::pair< std::string, int > > KTBN< GUM_SCALAR >::parents(std::string_view base,
357 int slice) const {
358 return _determineNodeSet_(_bn_.parents(_validateVariable_(base, slice)));
359 }
360
361 template < GUM_Numeric GUM_SCALAR >
362 std::vector< std::pair< std::string, int > >
363 KTBN< GUM_SCALAR >::parents(std::string_view node_name) const {
364 const auto [b, s] = _determineNode_(std::string{node_name});
365 return parents(b, s);
366 }
367
368 template < GUM_Numeric GUM_SCALAR >
369 std::vector< std::pair< std::string, int > > KTBN< GUM_SCALAR >::children(std::string_view base,
370 int slice) const {
371 return _determineNodeSet_(_bn_.children(_validateVariable_(base, slice)));
372 }
373
374 template < GUM_Numeric GUM_SCALAR >
375 std::vector< std::pair< std::string, int > >
376 KTBN< GUM_SCALAR >::children(std::string_view node_name) const {
377 const auto [b, s] = _determineNode_(std::string{node_name});
378 return children(b, s);
379 }
380
381 template < GUM_Numeric GUM_SCALAR >
382 void KTBN< GUM_SCALAR >::erase(std::string_view base) {
383 const std::string baseStr{base};
384
385 if (_temporal_.contains(baseStr)) {
386 for (Size t = 0; t < _k_; ++t)
387 _bn_.erase(_encode_(baseStr, static_cast< int >(t)));
388 _temporal_.erase(baseStr);
389 } else if (_atemporal_.contains(baseStr)) {
390 _bn_.erase(baseStr);
391 _atemporal_.erase(baseStr);
392 } else {
393 GUM_ERROR(NotFound, "No variable named '" << baseStr << "' in the k-DBN.")
394 }
395 }
396
397 template < GUM_Numeric GUM_SCALAR >
398 void KTBN< GUM_SCALAR >::changeVariableName(std::string_view oldBase, std::string_view newBase) {
399 const std::string oldStr{oldBase};
400 const std::string newStr{newBase};
401
402 if (oldStr == newStr) return;
403 if (newStr.empty()) GUM_ERROR(InvalidArgument, "New base name must not be empty.")
404
405 if (!_temporal_.contains(oldStr) && !_atemporal_.contains(oldStr))
406 GUM_ERROR(NotFound, "No variable named '" << oldStr << "' in the k-DBN.")
407 if (_temporal_.contains(newStr) || _atemporal_.contains(newStr))
408 GUM_ERROR(DuplicateLabel, "A variable named '" << newStr << "' already exists.")
409
410 if (_temporal_.contains(oldStr)) {
411 for (Size t = 0; t < _k_; ++t)
412 if (_atemporal_.contains(_encode_(newStr, static_cast< int >(t))))
414 "Renaming to '" << newStr << "': slice " << t
415 << " collides with an atemporal variable.")
416 for (Size t = 0; t < _k_; ++t)
417 _bn_.changeVariableName(_encode_(oldStr, static_cast< int >(t)),
418 _encode_(newStr, static_cast< int >(t)));
419 _temporal_.erase(oldStr);
420 _temporal_.insert(newStr);
421 } else {
422 const auto [decodedBase, decodedSlice] = _decodeName_(newStr);
423 if (decodedSlice != ATEMPORAL && _temporal_.contains(decodedBase))
425 "'" << newStr << "' conflicts with temporal process '" << decodedBase
426 << "': that name is already used by its slice nodes.")
427 _bn_.changeVariableName(oldStr, newStr);
428 _atemporal_.erase(oldStr);
429 _atemporal_.insert(newStr);
430 }
431 }
432
433 template < GUM_Numeric GUM_SCALAR >
434 NodeId KTBN< GUM_SCALAR >::_validateVariable_(std::string_view base, int slice) const {
435 const std::string baseStr{base};
436
437 if (slice == ATEMPORAL) {
438 if (!_atemporal_.contains(baseStr)) {
439 if (_temporal_.contains(baseStr))
441 "'" << baseStr << "' is a temporal process but is used as atemporal.")
442 GUM_ERROR(NotFound, "There is no atemporal variable named '" << baseStr << "'.")
443 }
444 return _bn_.idFromName(baseStr);
445 }
446
447 if (!_temporal_.contains(baseStr)) {
448 if (_atemporal_.contains(baseStr))
450 "'" << baseStr << "' is an atemporal variable but is used at slice " << slice
451 << ".")
452 GUM_ERROR(NotFound, "There is no temporal process named '" << baseStr << "'.")
453 }
454 if (slice < 0 || Size(slice) >= _k_)
456 "Slice " << slice << " is out of [0," << (_k_ - 1) << "] for process '" << baseStr
457 << "'.")
458 return _bn_.idFromName(_encode_(baseStr, slice));
459 }
460
461 template < GUM_Numeric GUM_SCALAR >
462 INLINE const DiscreteVariable& KTBN< GUM_SCALAR >::variable(std::string_view base,
463 int slice) const {
464 return _bn_.variable(_validateVariable_(base, slice));
465 }
466
467 template < GUM_Numeric GUM_SCALAR >
468 INLINE const DiscreteVariable& KTBN< GUM_SCALAR >::variable(std::string_view node_name) const {
469 const auto [b, s] = _determineNode_(std::string{node_name});
470 return variable(b, s);
471 }
472
473 template < GUM_Numeric GUM_SCALAR >
474 INLINE int KTBN< GUM_SCALAR >::timeSlice(const DiscreteVariable& var) const {
475 _bn_.idFromName(var.name()); // throws NotFound if var is not in this k-DBN
476 return _determineNode_(var.name()).second;
477 }
478
479 template < GUM_Numeric GUM_SCALAR >
480 INLINE std::string KTBN< GUM_SCALAR >::baseName(const DiscreteVariable& var) const {
481 _bn_.idFromName(var.name()); // throws NotFound if var is not in this k-DBN
482 return _determineNode_(var.name()).first;
483 }
484
485 // ===========================================================================
486 // Arc management
487 // ===========================================================================
488
489 template < GUM_Numeric GUM_SCALAR >
490 void KTBN< GUM_SCALAR >::addArc(std::string_view tailBase,
491 int tailSlice,
492 std::string_view headBase,
493 int headSlice) {
494 const NodeId tail = _validateVariable_(tailBase, tailSlice);
495 const NodeId head = _validateVariable_(headBase, headSlice);
496
497 if (headSlice == ATEMPORAL) {
498 if (tailSlice != ATEMPORAL)
500 "A temporal variable cannot be a parent of the atemporal variable '" << headBase
501 << "'.")
502 } else if (tailSlice != ATEMPORAL && tailSlice > headSlice) {
504 "An arc cannot go from a future slice (" << tailSlice << ") to a past slice ("
505 << headSlice << ").")
506 }
507
508 _bn_.addArc(tail, head);
509 }
510
511 template < GUM_Numeric GUM_SCALAR >
512 void KTBN< GUM_SCALAR >::eraseArc(std::string_view tailBase,
513 int tailSlice,
514 std::string_view headBase,
515 int headSlice) {
516 _bn_.eraseArc(_validateVariable_(tailBase, tailSlice), _validateVariable_(headBase, headSlice));
517 }
518
519 template < GUM_Numeric GUM_SCALAR >
520 bool KTBN< GUM_SCALAR >::existsArc(std::string_view tailBase,
521 int tailSlice,
522 std::string_view headBase,
523 int headSlice) const {
524 return _bn_.existsArc(_validateVariable_(tailBase, tailSlice),
525 _validateVariable_(headBase, headSlice));
526 }
527
528 template < GUM_Numeric GUM_SCALAR >
529 void KTBN< GUM_SCALAR >::addArc(std::string_view tail, std::string_view head) {
530 const auto [tb, ts] = _determineNode_(std::string{tail});
531 const auto [hb, hs] = _determineNode_(std::string{head});
532 addArc(tb, ts, hb, hs);
533 }
534
535 template < GUM_Numeric GUM_SCALAR >
536 void KTBN< GUM_SCALAR >::eraseArc(std::string_view tail, std::string_view head) {
537 const auto [tb, ts] = _determineNode_(std::string{tail});
538 const auto [hb, hs] = _determineNode_(std::string{head});
539 eraseArc(tb, ts, hb, hs);
540 }
541
542 template < GUM_Numeric GUM_SCALAR >
543 bool KTBN< GUM_SCALAR >::existsArc(std::string_view tail, std::string_view head) const {
544 const auto [tb, ts] = _determineNode_(std::string{tail});
545 const auto [hb, hs] = _determineNode_(std::string{head});
546 return existsArc(tb, ts, hb, hs);
547 }
548
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()) {
555 result.emplace_back(_determineNode_(_bn_.variable(arc.tail()).name()),
556 _determineNode_(_bn_.variable(arc.head()).name()));
557 }
558 return result;
559 }
560
561 // ===========================================================================
562 // Conditional probability tables
563 // ===========================================================================
564
565 template < GUM_Numeric GUM_SCALAR >
566 INLINE const Tensor< GUM_SCALAR >& KTBN< GUM_SCALAR >::cpt(std::string_view base,
567 int slice) const {
568 return _bn_.cpt(_validateVariable_(base, slice));
569 }
570
571 template < GUM_Numeric GUM_SCALAR >
572 INLINE const Tensor< GUM_SCALAR >& KTBN< GUM_SCALAR >::cpt(std::string_view node_name) const {
573 const auto [b, s] = _determineNode_(std::string{node_name});
574 return cpt(b, s);
575 }
576
577 template < GUM_Numeric GUM_SCALAR >
579 _bn_.generateCPTs();
580 }
581
582 template < GUM_Numeric GUM_SCALAR >
583 INLINE void KTBN< GUM_SCALAR >::generateCPT(std::string_view base, int slice) const {
584 _bn_.generateCPT(_validateVariable_(base, slice));
585 }
586
587 template < GUM_Numeric GUM_SCALAR >
588 INLINE void KTBN< GUM_SCALAR >::generateCPT(std::string_view node_name) const {
589 const auto [b, s] = _determineNode_(std::string{node_name});
590 generateCPT(b, s);
591 }
592
593 template < GUM_Numeric GUM_SCALAR >
595 std::string_view base,
596 int slice,
597 const std::map< std::pair< std::string, int >, KTBNModality >& parents,
598 const std::vector< GUM_SCALAR >& distribution) const {
599 const NodeId id = _validateVariable_(base, slice);
600 const Tensor< GUM_SCALAR >& cpt = _bn_.cpt(id);
601 const DiscreteVariable& self = _bn_.variable(id);
602
603 if (distribution.size() != self.domainSize())
605 "fillCPT: distribution has " << distribution.size() << " value(s) but '" << base
606 << "' has " << self.domainSize() << " modalities.")
607
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.")
613
614 // Address each parent by its (base, slice) identity — order-independent.
615 // No duplicate-parent check needed: a dictionary key is unique by
616 // construction, and here each (base, slice) pair names exactly one node,
617 // so no two distinct keys can alias the same parent.
618 Instantiation inst(cpt);
619 for (const auto& [parNode, parVal]: parents) {
620 const auto& [parBase, parSlice] = parNode;
621 const NodeId parId = _validateVariable_(parBase, parSlice);
622 const DiscreteVariable& parVar = _bn_.variable(parId);
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));
627 }
628
629 // Write the whole conditional distribution over the node's own modalities.
630 for (Idx m = 0; m < self.domainSize(); ++m) {
631 inst.chgVal(self, m);
632 cpt.set(inst, distribution[m]);
633 }
634 }
635
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 >&
640 parents,
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);
644 const DiscreteVariable& self = _bn_.variable(id);
645
646 if (distribution.size() != self.domainSize())
648 "fillCPT: distribution has " << distribution.size() << " value(s) but '"
649 << node_name << "' has " << self.domainSize()
650 << " modalities.")
651
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.")
657
658 // Unlike the (base, slice)-keyed overload above, duplicates ARE possible here
659 // despite unique map keys: "X[0]" and (base="X", slice=0) compare unequal as
660 // std::variant values yet name the same node. Both are resolved to an engine
661 // name first, so the seen-check below can catch the alias.
662 Instantiation inst(cpt);
663 NodeSet seen;
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);
671 const DiscreteVariable& parVar = _bn_.variable(parId);
672 if (parId == id || !cpt.contains(parVar))
674 "fillCPT: '" << parName << "' is not a parent of the target node.")
675 if (seen.contains(parId))
676 GUM_ERROR(InvalidArgument, "fillCPT: parent '" << parName << "' is listed more than once.")
677 seen.insert(parId);
678 inst.chgVal(parVar, parVal.toIndex(parVar));
679 }
680
681 for (Idx m = 0; m < self.domainSize(); ++m) {
682 inst.chgVal(self, m);
683 cpt.set(inst, distribution[m]);
684 }
685 }
686
687 // ===========================================================================
688 // Transformations
689 // ===========================================================================
690
691 template < GUM_Numeric GUM_SCALAR >
692 BayesNet< GUM_SCALAR > KTBN< GUM_SCALAR >::toBN() const {
693 return BayesNet< GUM_SCALAR >(_bn_);
694 }
695
696 template < GUM_Numeric GUM_SCALAR >
697 BayesNet< GUM_SCALAR > KTBN< GUM_SCALAR >::unroll(Size nbTimeSlices) const {
698 if (nbTimeSlices < _k_)
700 "Cannot unroll over " << nbTimeSlices << " slices: fewer than the order k=" << _k_
701 << ".")
702
703 BayesNet< GUM_SCALAR > unrolled;
704 const int kernelSlice = static_cast< int >(_k_ - 1);
705
706 // 1. atemporal variables (kept as-is)
707 for (const auto& a: _atemporal_) {
708 unrolled.add(_bn_.variable(_bn_.idFromName(a)));
709 }
710
711 // 2. temporal variables, instantiated for every slice 0..nbTimeSlices-1
712 for (const auto& p: _temporal_) {
713 const DiscreteVariable& templateVar
714 = _bn_.variable(_bn_.idFromName(_encode_(p, kernelSlice)));
715 for (Size t = 0; t < nbTimeSlices; ++t) {
716 // clone() to rename; BayesNet::add() clones again internally (unavoidable via public API).
717 std::unique_ptr< DiscreteVariable > clone(templateVar.clone());
718 clone->setName(_encode_(p, static_cast< int >(t)));
719 unrolled.add(*clone);
720 }
721 }
722
723 // 3. template arcs (slices 0..k-1) are copied verbatim: the engine names of
724 // the template already match the unrolled names for those slices.
725 for (const auto& arc: _bn_.arcs()) {
726 unrolled.addArc(_bn_.variable(arc.tail()).name(), _bn_.variable(arc.head()).name());
727 }
728
729 // 4. CPTs of the template nodes (slices 0..k-1): copied by name.
730 for (const NodeId n: _bn_.nodes()) {
731 unrolled.cpt(_bn_.variable(n).name()).fillWith(_bn_.cpt(n));
732 }
733
734 // 5. transition kernel: for each extra slice t = k..nbTimeSlices-1,
735 // add arcs and fill the CPT in one pass using lags computed once per process.
736 for (const auto& p: _temporal_) {
737 const NodeId lastSliceNodeId = _bn_.idFromName(_encode_(p, kernelSlice));
738 const Tensor< GUM_SCALAR >& templateCpt = _bn_.cpt(lastSliceNodeId);
739
740 // (parBase, lag): lag == ATEMPORAL for static parents, otherwise lag = (k-1) - parSlice.
741 std::vector< std::pair< std::string, int > > lags;
742 for (const auto& [parBase, parSlice]: parents(p, kernelSlice)) {
743 const int lag = (parSlice == ATEMPORAL) ? ATEMPORAL : kernelSlice - parSlice;
744 lags.emplace_back(parBase, lag);
745 }
746
747 HashTable< std::string, std::string > unrolledToTemplate;
748 std::vector< std::string > templateVarNames;
749
750 for (Size t = _k_; t < nbTimeSlices; ++t) {
751 const std::string child = _encode_(p, static_cast< int >(t));
752
753 // Add arcs and build the unrolled->template name mapping simultaneously.
754 unrolledToTemplate.clear();
755 unrolledToTemplate.insert(child, _encode_(p, kernelSlice));
756 for (const auto& [parBase, lag]: lags) {
757 if (lag == ATEMPORAL) {
758 unrolled.addArc(parBase, child);
759 unrolledToTemplate.insert(parBase, parBase);
760 } else {
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));
764 }
765 }
766
767 // Fill the CPT using the mapping built above.
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()]);
773 }
774 unrolledCpt.fillWith(templateCpt, templateVarNames);
775 }
776 }
777
778 return unrolled;
779 }
780
781 // ===========================================================================
782 // Persistence and conversion
783 // ===========================================================================
784
785 template < GUM_Numeric GUM_SCALAR >
786 std::pair< std::string, bool > KTBN< GUM_SCALAR >::_resolveGumFormat_(std::string_view filename) {
787 // The extension selects the format: ".jgum" is text, anything else is binary
788 // and gets a ".bgum" extension appended if missing.
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}; // .second = binary
793 }
794
795 template < GUM_Numeric GUM_SCALAR >
796 void KTBN< GUM_SCALAR >::save(std::string_view filename) const {
797 const auto [filepath, binary] = _resolveGumFormat_(filename);
798
799 // Persist the temporal/atemporal classification (and k) as BN properties so that
800 // load() restores the k-DBN exactly, instead of re-deriving it heuristically from
801 // node names (ambiguous for k=1 processes and bracket-named atemporal variables).
802 // Separator is ','; '\' and ',' inside names are backslash-escaped so any name is safe.
803 auto join = [](const std::unordered_set< std::string >& names) {
804 std::string out;
805 for (const auto& n: names) {
806 if (!out.empty()) out += ',';
807 for (const char c: n) {
808 if (c == '\\' || c == ',') out += '\\';
809 out += c;
810 }
811 }
812 return out;
813 };
814
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_));
819 GumBNWriter< GUM_SCALAR > writer(binary);
820 writer.write(filepath, annotated);
821 }
822
823 template < GUM_Numeric GUM_SCALAR >
824 KTBN< GUM_SCALAR > KTBN< GUM_SCALAR >::load(std::string_view filename) {
825 const auto [filepath, binary] = _resolveGumFormat_(filename);
826
827 BayesNet< GUM_SCALAR > bn;
828 GumBNReader< GUM_SCALAR > reader(&bn, filepath, binary);
829 const Size nbErr = reader.proceed();
830 if (nbErr > 0) {
831 std::stringstream stream;
832 reader.showElegantErrorsAndWarnings(stream);
833 reader.showErrorCounts(stream);
834 GUM_ERROR(IOError, "KTBN::load: " << stream.str())
835 }
836
837 // A file written by save() carries the classification as properties: restore it
838 // directly. Otherwise fall back to fromBN(), which re-derives it from node names.
839 if (!(bn.existsProperty("KTBN.k") && bn.existsProperty("KTBN.temporal")
840 && bn.existsProperty("KTBN.atemporal")))
841 return fromBN(bn);
842
843 auto split = [](std::string_view csv, std::unordered_set< std::string >& out) {
844 std::string current;
845 bool escaped = false;
846 for (const char c: csv) {
847 if (escaped) {
848 current += c;
849 escaped = false;
850 } else if (c == '\\') {
851 escaped = true;
852 } else if (c == ',') {
853 if (!current.empty()) out.insert(current);
854 current.clear();
855 } else {
856 current += c;
857 }
858 }
859 if (!current.empty()) out.insert(current);
860 };
861
862 Size k_val{};
863 try {
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())
869 }
870 KTBN< GUM_SCALAR > res(k_val);
871 res._bn_ = bn;
872 split(bn.property("KTBN.temporal"), res._temporal_);
873 split(bn.property("KTBN.atemporal"), res._atemporal_);
874
875 return res;
876 }
877
878 template < GUM_Numeric GUM_SCALAR >
879 KTBN< GUM_SCALAR >
880 KTBN< GUM_SCALAR >::fromBN(const BayesNet< GUM_SCALAR >& bn,
881 const std::unordered_set< std::string >& atemporalNodes,
882 std::vector< std::string >* warnings) {
883 KTBN< GUM_SCALAR > res(1); // _determineNodesFromBN_ below will modify this k=1
884 res._bn_ = bn;
885 res._determineNodesFromBN_(atemporalNodes, warnings);
886 return res;
887 }
888
889 template < GUM_Numeric GUM_SCALAR >
891 const std::unordered_set< std::string >& atemporalNodes,
892 std::vector< std::string >* warnings) {
893 _temporal_.clear();
894 _atemporal_.clear();
895
896 // a declared name must be a node of the BN: checked before any mutation, so
897 // a typo cannot leave the object half-built
898 for (const std::string& name: atemporalNodes)
899 if (!_bn_.exists(name))
900 GUM_ERROR(NotFound, "fromBN: '" << name << "' is not a node of the BN.")
901
902 const auto warn = [warnings](const std::string& message) {
903 if (warnings != nullptr) warnings->push_back(message);
904 };
905
906 // fromBN()'s bracket-free convention: a name ending in a run of digits
907 // denotes a temporal variable at the timeslice given by that integer (base =
908 // everything before the run); a name with no trailing digit is atemporal.
909 // Purely syntactic, like _decodeName_, but that one expects the engine's own
910 // base[t] convention. Used below only when the source BN carries no bracket
911 // at all (see hasBracket).
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])))
915 --pos;
916 if (pos == name.size()) return {std::string{name}, ATEMPORAL};
917
918 const std::string_view digits = name.substr(pos);
919 int slice{};
920 try {
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.")
925 }
926 return {std::string{name.substr(0, pos)}, slice};
927 };
928
929 // The two conventions never mix within one graph: if any node name already
930 // carries the engine's own base[t] bracket notation, the WHOLE graph is read
931 // that way (legacy behaviour -- what KTBNLearner and toBN() round-trips
932 // produce); only when no node carries a bracket at all does every node get
933 // read via the trailing-integer convention above. Nodes declared in
934 // atemporalNodes are skipped here: their shape says nothing about the rest
935 // of the graph's convention, so a bracket-shaped one (e.g. an orphan
936 // "Y[0]" named atemporal on purpose, see below) must not force
937 // bracket-reading onto otherwise bracket-free temporal nodes.
938 bool hasBracket = false;
939 for (const NodeId n: _bn_.nodes()) {
940 const std::string& name = _bn_.variable(n).name();
941 if (atemporalNodes.contains(name)) continue;
942 if (_decodeName_(name).second != ATEMPORAL) {
943 hasBracket = true;
944 break;
945 }
946 }
947
948 // first pass: determine every node, collecting the slices seen per process
949 std::vector< std::string > discovered; // temporal bases, in order
951 int maxSlice = -1;
952
953 for (const NodeId n: _bn_.nodes()) {
954 const std::string& name = _bn_.variable(n).name();
955
956 // an explicitly declared node is atemporal whatever its shape, trailing
957 // digits included: it never enters slicesPerProcess, so the completeness
958 // rules below never see it and never warn about it
959 if (atemporalNodes.contains(name)) {
960 _atemporal_.insert(name);
961 continue;
962 }
963
964 const auto [base, slice] = hasBracket ? _decodeName_(name) : decodeTrailingSlice(name);
965
966 if (slice == ATEMPORAL) {
967 _atemporal_.insert(base);
968 } else {
969 if (!slicesPerProcess.exists(base)) {
970 slicesPerProcess.insert(base, HashTable< int, NodeId >());
971 discovered.push_back(base);
972 }
973
974 if (slicesPerProcess[base].exists(slice))
976 "Two variables map to process '" << base << "' at slice " << slice << ".")
977
978 slicesPerProcess[base].insert(slice, n);
979
980 if (slice > maxSlice) maxSlice = slice;
981 }
982 }
983
984 _k_ = (maxSlice < 0) ? Size(1) : Size(maxSlice + 1);
985
986 // second pass: register the complete processes. A group that does not cover
987 // every slice 0..k-1 is NOT rejected: each of its nodes becomes an atemporal
988 // variable, original name kept, and a warning is recorded. Only a base used
989 // BOTH bare and temporal-shaped stays an error -- there the two readings
990 // collide on one name, and no reclassification can resolve that.
991 const char* const conventionNoun = hasBracket ? "bracket" : "digit-suffixed";
992 for (const auto& base: discovered) {
993 const HashTable< int, NodeId >& sliceMap = slicesPerProcess[base];
994
995 if (_atemporal_.contains(base))
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().")
1000
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);
1006 }
1007
1008 if (missing.empty()) {
1009 // Every slice of a process must carry the SAME variable. add() cannot
1010 // break this -- it clones one variable into k instances -- so fromBN()
1011 // is the only way in, and nothing downstream re-checks: KTBNInference
1012 // sizes every ring slot of a process from the kernel slice alone, so a
1013 // slice with a divergent domain would be triangulated against the wrong
1014 // size and then filled with a tensor of another dimension.
1015 // Names necessarily differ between slices, and Variable::operator==
1016 // compares them, so the reference is renamed onto each slice in turn.
1017 const DiscreteVariable& ref = _bn_.variable(sliceMap[0]);
1018 std::unique_ptr< DiscreteVariable > probe(ref.clone());
1019 for (Size t = 1; t < _k_; ++t) {
1020 const DiscreteVariable& other = _bn_.variable(sliceMap[static_cast< int >(t)]);
1021 probe->setName(other.name());
1022 if (!(*probe == other))
1024 "The temporal process '"
1025 << base << "' has mismatched slice variables: '" << ref.name() << "' is "
1026 << ref.domain() << " but '" << other.name() << "' is " << other.domain()
1027 << ". Every slice of a process must have the same type and domain.")
1028 }
1029 _temporal_.insert(base);
1030 if (!hasBracket) {
1031 // Under the trailing-integer convention the slices just matched above
1032 // carry no bracket notation yet: rename them onto the engine's
1033 // canonical base[t] form -- the invariant every other method (add(),
1034 // unroll(), rename(), ...) relies on. Under the bracket convention
1035 // source names are already canonical, so nothing to do here.
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)));
1039 }
1040 continue;
1041 }
1042
1043 std::string reclassified;
1044 for (auto it = sliceMap.cbegin(); it != sliceMap.cend(); ++it) {
1045 const std::string& nodeName = _bn_.variable(it.val()).name();
1046 _atemporal_.insert(nodeName);
1047 if (!reclassified.empty()) reclassified += ", ";
1048 reclassified += "'" + nodeName + "'";
1049 }
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.");
1055 }
1056
1057 // Every surviving process holds exactly the slices 0..k-1, so k is still the
1058 // one the first pass computed -- unless none survived, in which case the
1059 // largest slice index was contributed by a group that is now atemporal and
1060 // k has nothing left to describe.
1061 if (_temporal_.empty()) _k_ = Size(1);
1062
1063
1064 // third pass: validate temporal causality of the foreign arcs.
1065 for (const auto& arc: _bn_.arcs()) {
1066 const int tailSlice = _determineNode_(_bn_.variable(arc.tail()).name()).second;
1067 const int headSlice = _determineNode_(_bn_.variable(arc.head()).name()).second;
1068 if (headSlice == ATEMPORAL && tailSlice != ATEMPORAL)
1070 "The network has a temporal->atemporal arc into '"
1071 << _bn_.variable(arc.head()).name() << "'.")
1072 if (headSlice != ATEMPORAL && tailSlice != ATEMPORAL && tailSlice > headSlice)
1074 "The network has a future->past arc " << _bn_.variable(arc.tail()).name() << "->"
1075 << _bn_.variable(arc.head()).name() << ".")
1076 }
1077 }
1078
1079 // ===========================================================================
1080 // Various
1081 // ===========================================================================
1082
1083 template < GUM_Numeric GUM_SCALAR >
1084 std::string KTBN< GUM_SCALAR >::toString() const {
1085 const auto join = [](const std::vector< std::string >& v) {
1086 std::string s;
1087 for (const auto& n: v) {
1088 if (!s.empty()) s += ", ";
1089 s += n;
1090 }
1091 return s;
1092 };
1093
1094 const std::vector< std::string > temporalNames(_temporal_.begin(), _temporal_.end());
1095 const std::vector< std::string > atemporalNames(_atemporal_.begin(), _atemporal_.end());
1096
1097 std::string arcs;
1098 for (const auto& arc: _bn_.arcs())
1099 arcs += std::format(" {} -> {}\n",
1100 _bn_.variable(arc.tail()).name(),
1101 _bn_.variable(arc.head()).name());
1102
1103 return std::format("k-TBN (k={}, {} nodes, {} arcs)\n"
1104 " temporal processes ({}): {}\n"
1105 " atemporal variables ({}): {}\n"
1106 "\n arcs ({}):\n{}",
1107 _k_,
1108 _bn_.size(),
1109 _bn_.sizeArcs(),
1110 _temporal_.size(),
1111 join(temporalNames),
1112 _atemporal_.size(),
1113 join(atemporalNames),
1114 _bn_.sizeArcs(),
1115 arcs);
1116 }
1117
1118 template < GUM_Numeric GUM_SCALAR >
1119 std::string KTBN< GUM_SCALAR >::toDot() const {
1120 return _timeSlicesToDot_(_bn_, false);
1121 }
1122
1123 template < GUM_Numeric GUM_SCALAR >
1124 std::string KTBN< GUM_SCALAR >::toUnrolledDot(Size T, bool highlightReplicated) const {
1125 if (T < _k_)
1126 GUM_ERROR(OperationNotAllowed, "toUnrolledDot: T=" << T << " must be >= k=" << _k_ << ".")
1127 return _timeSlicesToDot_(unroll(T), highlightReplicated);
1128 }
1129
1130 template < GUM_Numeric GUM_SCALAR >
1131 std::string KTBN< GUM_SCALAR >::_escapeDot_(std::string_view name) {
1132 std::string out;
1133 out.reserve(name.size());
1134 for (const char c: name) {
1135 if (c == '"') out += '\\';
1136 out += c;
1137 }
1138 return out;
1139 }
1140
1141 // Direct port of pyAgrum's pyagrum.lib.dynamicBN._TimeSlicesToDot: groups nodes by
1142 // timeslice (one cluster per slice, atemporal nodes ungrouped), draws the real arcs
1143 // unconstrained, then chains each temporal variable across consecutive slices with
1144 // invisible edges so every cluster keeps the same vertical variable order.
1145 template < GUM_Numeric GUM_SCALAR >
1146 std::string KTBN< GUM_SCALAR >::_timeSlicesToDot_(const BayesNet< GUM_SCALAR >& bn,
1147 bool highlightReplicated) const {
1148 // Group (full name, base label) by timeslice. std::map keeps keys sorted, and
1149 // ATEMPORAL == -1 so atemporal variables naturally sort first, followed by
1150 // increasing slice indices — mirroring pyAgrum's noTimeCluster-then-slices order.
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();
1154 const auto [base, slice] = _decodeName_(name);
1155 timeslices[slice].emplace_back(name, base);
1156 }
1157
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";
1163
1164 for (auto& [slice, nodes]: timeslices) {
1165 std::sort(nodes.begin(), nodes.end());
1166 if (slice == ATEMPORAL) {
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)
1172 dot << " \"" << _escapeDot_(full) << "\" [label=\"" << _escapeDot_(label) << "\"];\n";
1173 dot << " }\n";
1174 } else {
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)
1181 dot << " \"" << _escapeDot_(full) << "\" [label=\"" << _escapeDot_(label) << "\"];\n";
1182 dot << " }\n";
1183 }
1184 dot << "\n";
1185 }
1186
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";
1191
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;
1196 int prec = ATEMPORAL;
1197 bool first = true;
1198 for (const auto& [slice, nodes]: timeslices) {
1199 if (slice == ATEMPORAL) continue;
1200 if (!first)
1201 dot << " \"" << _escapeDot_(_encode_(label, prec)) << "\" -> \""
1202 << _escapeDot_(_encode_(label, slice)) << "\";\n";
1203 prec = slice;
1204 first = false;
1205 }
1206 }
1207 }
1208
1209 dot << "}\n";
1210 return dot.str();
1211 }
1212
1213 template < GUM_Numeric GUM_SCALAR >
1214 INLINE std::string KTBN< GUM_SCALAR >::bnToDot() const {
1215 return _bn_.toDot();
1216 }
1217
1218 template < GUM_Numeric GUM_SCALAR >
1220 std::set< std::string > baseNames(_temporal_.begin(), _temporal_.end());
1221 baseNames.insert(_atemporal_.begin(), _atemporal_.end());
1222
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; // not part of the repeated transition kernel
1229
1230 const int lag = (tailSlice == ATEMPORAL) ? ATEMPORAL : headSlice - tailSlice;
1231 edges.emplace(tailBase, headBase, lag);
1232 }
1233
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";
1238
1239 for (const auto& base: baseNames)
1240 dot << " \"" << _escapeDot_(base) << "\";\n";
1241 dot << "\n";
1242
1243 for (const auto& [tailBase, headBase, lag]: edges) {
1244 dot << " \"" << _escapeDot_(tailBase) << "\" -> \"" << _escapeDot_(headBase) << "\"";
1245 if (lag != ATEMPORAL) dot << " [label=\"" << lag << "\"]";
1246 dot << ";\n";
1247 }
1248
1249 dot << "}\n";
1250 return dot.str();
1251 }
1252
1253 template < GUM_Numeric GUM_SCALAR >
1254 std::ostream& operator<<(std::ostream& output, const KTBN< GUM_SCALAR >& kdbn) {
1255 output << kdbn.toString();
1256 return output;
1257 }
1258
1259} // namespace gum
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.
Definition GumBNWriter.h:80
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).
Definition KTBN_tpl.h:369
void clear()
Removes all variables and arcs, keeping the order .
Definition KTBN_tpl.h:204
void addTemporal(const DiscreteVariable &var)
Convenience shortcut for add(var, true).
Definition KTBN_tpl.h:292
const std::unordered_set< std::string > & temporalVarNames() const
Definition KTBN_tpl.h:324
void addAtemporal(const DiscreteVariable &var)
Convenience shortcut for add(var, false).
Definition KTBN_tpl.h:297
BayesNet< GUM_SCALAR > toBN() const
Definition KTBN_tpl.h:692
void addArc(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice)
Adds an arc between two (process, slice) endpoints.
Definition KTBN_tpl.h:490
bool empty() const
Definition KTBN_tpl.h:199
std::vector< std::pair< std::pair< std::string, int >, std::pair< std::string, int > > > arcs() const
Definition KTBN_tpl.h:551
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...
Definition KTBN_tpl.h:84
std::vector< std::pair< std::string, int > > nodes() const
Definition KTBN_tpl.h:344
BayesNet< GUM_SCALAR > _bn_
The underlying Bayesian network used as a storage engine for the template.
Definition KTBN.h:707
static KTBN< GUM_SCALAR > load(std::string_view filename)
Loads a k-DBN from a GUM file produced by save().
Definition KTBN_tpl.h:824
void erase(std::string_view base)
Removes a variable and all its incident arcs.
Definition KTBN_tpl.h:382
const DiscreteVariable & variable(std::string_view base, int slice) const
Returns the gum::DiscreteVariable of a (process, slice) couple.
Definition KTBN_tpl.h:462
std::string summaryGraph() const
Returns the Graphviz DOT string of the summary graph: the projection of the transition kernel alone (...
Definition KTBN_tpl.h:1219
bool existsArc(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice) const
Definition KTBN_tpl.h:520
void save(std::string_view filename) const
Saves the template in the GUM format (text .jgum or binary .bgum).
Definition KTBN_tpl.h:796
static std::string _escapeDot_(std::string_view name)
Escapes double quotes for a DOT identifier or label. Shared by timeSlicesToDot() and summaryGraph().
Definition KTBN_tpl.h:1131
NodeId _validateVariable_(std::string_view base, int slice) const
Resolves and validates a (base, slice) endpoint into its NodeId.
Definition KTBN_tpl.h:434
std::unordered_set< std::string > _temporal_
Base names of the registered temporal processes.
Definition KTBN.h:710
void eraseArc(std::string_view tailBase, int tailSlice, std::string_view headBase, int headSlice)
Removes an arc between two (process, slice) endpoints.
Definition KTBN_tpl.h:512
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).
Definition KTBN_tpl.h:356
const std::unordered_set< std::string > & atemporalVarNames() const
Definition KTBN_tpl.h:329
void generateCPT(std::string_view base, int slice) const
Randomly generates the CPT of a single node.
Definition KTBN_tpl.h:583
Size nbAtemporalVars() const
Definition KTBN_tpl.h:339
std::string _encode_(std::string_view base, int slice) const
Encodes (base, slice) → engine name: base[t], or base if atemporal.
Definition KTBN_tpl.h:78
const Tensor< GUM_SCALAR > & cpt(std::string_view base, int slice) const
Returns the CPT of a (process, slice) couple.
Definition KTBN_tpl.h:566
void generateCPTs() const
Randomly generates every CPT of the template.
Definition KTBN_tpl.h:578
std::unordered_set< std::string > _atemporal_
Base names of the registered atemporal variables.
Definition KTBN.h:713
std::string baseName(const DiscreteVariable &var) const
Returns the base name (without bracket encoding) of var.
Definition KTBN_tpl.h:480
static constexpr int ATEMPORAL
Conventional time-slice value denoting an atemporal (static) variable.
Definition KTBN.h:200
Size sizeArcs() const
Definition KTBN_tpl.h:194
std::string toUnrolledDot(Size T, bool highlightReplicated=false) const
Returns a Graphviz DOT string of the k-DBN unrolled over T time slices.
Definition KTBN_tpl.h:1124
std::string bnToDot() const
Returns the Graphviz DOT string of the underlying storage BayesNet.
Definition KTBN_tpl.h:1214
KTBN< GUM_SCALAR > & operator=(const KTBN< GUM_SCALAR > &source)
Copy assignment operator.
Definition KTBN_tpl.h:156
int timeSlice(const DiscreteVariable &var) const
The time slice of var, or ATEMPORAL if it is atemporal.
Definition KTBN_tpl.h:474
void changeVariableName(std::string_view oldBase, std::string_view newBase)
Renames a variable (temporal process or atemporal variable).
Definition KTBN_tpl.h:398
std::string toDot() const
Returns a Graphviz DOT string with one cluster per time slice.
Definition KTBN_tpl.h:1119
virtual ~KTBN()
Destructor.
Definition KTBN_tpl.h:137
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...
Definition KTBN_tpl.h:109
BayesNet< GUM_SCALAR > unroll(Size nbTimeSlices) const
Unrolls the k-DBN into a standard gum::BayesNet.
Definition KTBN_tpl.h:697
KTBN(Size k=2)
Default constructor.
Definition KTBN_tpl.h:131
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).
Definition KTBN_tpl.h:594
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...
Definition KTBN_tpl.h:1146
Size _k_
The order (number of time slices in the template).
Definition KTBN.h:704
Size nbTemporalVars() const
Definition KTBN_tpl.h:334
std::vector< std::pair< std::string, int > > _determineNodeSet_(const NodeSet &ids) const
Maps a set of node ids to (base, slice) pairs (via determineNode).
Definition KTBN_tpl.h:118
bool exists(std::string_view base) const
Definition KTBN_tpl.h:318
void _validateAdd_(const std::string &base, bool temporal) const
Checks that a variable named base can be added.
Definition KTBN_tpl.h:215
Size size() const
Definition KTBN_tpl.h:189
std::string toString() const
Definition KTBN_tpl.h:1084
void add(const DiscreteVariable &var, bool temporal=true)
Adds a variable to the k-DBN.
Definition KTBN_tpl.h:264
static std::pair< std::string, bool > _resolveGumFormat_(std::string_view filename)
Resolves a user filename to (filepath, binary): ensures a .jgum/.bgum extension (....
Definition KTBN_tpl.h:786
Size k() const
Definition KTBN_tpl.h:184
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...
Definition KTBN_tpl.h:880
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...
Definition KTBN_tpl.h:890
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.
Definition set_tpl.h:468
void insert(const Key &k)
Inserts a new element into the set.
Definition set_tpl.h:510
Size size() const noexcept
Returns the number of elements in the set.
Definition set_tpl.h:607
Exception : problem with size.
const std::string & name() const
returns the name of the variable
#define GUM_ERROR(type, msg)
Definition exceptions.h:76
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Definition types.h:74
Size Idx
Type for indexes.
Definition types.h:79
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
Definition agrum.h:46
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.
Definition KTBN.h:95
Idx index
The index, when isLabel is false.
Definition KTBN.h:118
bool isLabel
Whether the value was spelled as a label rather than an index.
Definition KTBN.h:116
KTBNModality(T modality)
From a modality index.
Definition KTBN_tpl.h:71