aGrUM 3.2.0
a C++ library for (probabilistic) graphical models
KTBNLearner_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
42#pragma once
43
44
50
51#include <filesystem>
52#include <fstream>
53#include <sstream>
54
58
59namespace gum::learning {
60
61 // =========================================================================
62 // Internal-learner fan-out helpers
63 // =========================================================================
64
65 template < GUM_Numeric GUM_SCALAR >
66 template < class F >
72
73 template < GUM_Numeric GUM_SCALAR >
74 template < class F >
76 std::string_view head,
77 F&& f) {
78 const int last = static_cast< int >(_prior_ktbn_.k()) - 1;
79 const int tailSlice = _determineNode_(std::string{tail}).second;
80 const int headSlice = _determineNode_(std::string{head}).second;
81 const bool bothAtemp = (tailSlice == KTBN< GUM_SCALAR >::ATEMPORAL
82 && headSlice == KTBN< GUM_SCALAR >::ATEMPORAL);
83 if (headSlice == last) f(*_transitionLearner_);
84 else if (bothAtemp && _atemporalLearner_) f(*_atemporalLearner_);
85 else if (tailSlice != last) f(*_initialLearner_);
86 // else: the head is at a past slice but the tail sits on the kernel one,
87 // i.e. a backward arc. No table can hold it -- the initial one has no
88 // slice k-1 column -- and the k-TBN forbids it anyway, so there is
89 // nothing to constrain. Silently nothing, as the old broadcast did
90 // through its "tailSlice != last && headSlice != last" guard.
91 }
92
93 template < GUM_Numeric GUM_SCALAR >
94 template < class F >
96 std::string_view headBase,
97 F&& f) const {
98 if (!_isKnownBase_(tailBase))
100 "unknown base variable '" << tailBase
101 << "': it is not one of this learner's variables")
102 if (!_isKnownBase_(headBase))
104 "unknown base variable '" << headBase
105 << "': it is not one of this learner's variables")
106
107 constexpr int AT = KTBN< GUM_SCALAR >::ATEMPORAL;
108 const int k = static_cast< int >(_prior_ktbn_.k());
109 const bool tailAtemp = _prior_ktbn_.atemporalVarNames().contains(std::string{tailBase});
110 const bool headAtemp = _prior_ktbn_.atemporalVarNames().contains(std::string{headBase});
111
112 if (tailAtemp && headAtemp) {
113 f(AT, AT);
114 } else if (tailAtemp) {
115 for (int hs = 0; hs < k; ++hs)
116 f(AT, hs);
117 } else if (headAtemp) {
118 // temporal tail -> atemporal head is already structurally impossible
119 } else {
120 for (int ts = 0; ts < k; ++ts)
121 for (int hs = ts; hs < k; ++hs)
122 f(ts, hs);
123 }
124 }
125
126 // =========================================================================
127 // Constructors / Destructors
128 // =========================================================================
129
130 template < GUM_Numeric GUM_SCALAR >
132 std::string_view csvBaseName,
134 Size k,
135 const std::unordered_set< std::string >& atemporalVars,
136 const std::vector< std::string >& missingSymbols,
137 bool induceTypes,
138 bool ignoreMissingSymbols) :
139 _ignoreMissingSymbols_(ignoreMissingSymbols), _prior_ktbn_(_buildPriorFromCSV_(dirPath,
140 csvBaseName,
141 k,
142 atemporalVars,
143 missingSymbols,
144 induceTypes)) {
145 // k >= 2 is already enforced by _buildPriorFromCSV_ during member initialisation
146 if (nbSamples == 0) GUM_ERROR(InvalidArgument, "KTBNLearner needs at least one sample")
147 GUM_CONSTRUCTOR(KTBNLearner)
148 try {
149 _build_(dirPath, csvBaseName, nbSamples, missingSymbols);
150 } catch (const gum::UnknownLabelInDatabase&) {
151 // Domains are inferred from trajectory 1 alone, so a variable whose full
152 // domain is absent there fails once a later trajectory shows a new label.
153 // Atemporal variables are the usual culprits: constant within a trajectory,
154 // they reveal only one value per file. Re-throw with a KTBN-specific hint.
155 GUM_DESTRUCTOR(KTBNLearner)
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.")
166 } catch (...) {
167 GUM_DESTRUCTOR(KTBNLearner)
168 throw;
169 }
170 }
171
172 template < GUM_Numeric GUM_SCALAR >
174 std::string_view csvBaseName,
176 Size k,
177 const std::vector< std::string >& missingSymbols,
178 bool induceTypes,
179 bool ignoreMissingSymbols) :
180 // _inferAtemporalVars_ checks k>=2 itself, before opening anything (see
181 // its declaration) — this initialiser-list call runs before the
182 // delegated-to constructor's own body, so that check cannot be left to
183 // _buildPriorFromCSV_ the way the explicit constructor leaves it.
184 // Delegates to the explicit-atemporalVars constructor for the rest
185 // (including the UnknownLabelInDatabase hint); _inferAtemporalVars_
186 // needs nbSamples, unlike _buildPriorFromCSV_, since one trajectory
187 // alone cannot show that a value stays constant.
188 KTBNLearner(dirPath,
189 csvBaseName,
190 nbSamples,
191 k,
192 _inferAtemporalVars_(dirPath, csvBaseName, nbSamples, k, missingSymbols),
193 missingSymbols,
194 induceTypes,
195 ignoreMissingSymbols) {}
196
197 template < GUM_Numeric GUM_SCALAR >
199 std::string_view csvBaseName,
201 Size k,
202 const BayesNet< GUM_SCALAR >& bn,
203 const std::unordered_set< std::string >& atemporalVars,
204 const std::vector< std::string >& missingSymbols,
205 bool ignoreMissingSymbols) :
206 _ignoreMissingSymbols_(ignoreMissingSymbols),
207 _prior_ktbn_(_buildPriorFromBN_(k, bn, atemporalVars)) {
208 // k >= 2 is already enforced by _buildPriorFromBN_ during member initialisation
209 if (nbSamples == 0) GUM_ERROR(InvalidArgument, "KTBNLearner needs at least one sample")
210 GUM_CONSTRUCTOR(KTBNLearner)
211 try {
212 _build_(dirPath, csvBaseName, nbSamples, missingSymbols);
213 } catch (...) {
214 GUM_DESTRUCTOR(KTBNLearner)
215 throw;
216 }
217 }
218
219 template < GUM_Numeric GUM_SCALAR >
223
224 // =========================================================================
225 // Main learning methods
226 // =========================================================================
227
228 template < GUM_Numeric GUM_SCALAR >
230 // Honour the possible-edge whitelist for atemporal->atemporal arcs: when a
231 // whitelist is active (>=1 possible edge with a temporal endpoint) but no
232 // atemporal->atemporal edge was whitelisted, the atemporal learner has an
233 // empty (hence unrestricted) list. Skip it entirely so no atemporal arc is
234 // produced; _assemble_ then leaves the atemporal variables as roots.
235 const bool suppressAtemporal
237 BayesNet< GUM_SCALAR > transitionBN = _transitionLearner_->learnBN();
238 BayesNet< GUM_SCALAR > initialBN = _initialLearner_->learnBN();
239 BayesNet< GUM_SCALAR > atemporalBN = (_atemporalLearner_ && !suppressAtemporal)
240 ? _atemporalLearner_->learnBN()
241 : BayesNet< GUM_SCALAR >{};
242 return _assemble_(transitionBN, initialBN, atemporalBN);
243 }
244
245 template < GUM_Numeric GUM_SCALAR >
246 KTBN< GUM_SCALAR > KTBNLearner< GUM_SCALAR >::learnParameters(const KTBN< GUM_SCALAR >& structure,
247 bool takeIntoAccountScore) {
248 if (structure.k() != _prior_ktbn_.k())
250 "learnParameters: structure has k="
251 << structure.k() << " but this learner was built with k=" << _prior_ktbn_.k())
252 const int k = (int)_prior_ktbn_.k();
253
254 // Build one DAG per internal table, in that table's own NodeId space, by
255 // routing each of structure's (base, slice) arcs to the table that owns its
256 // head: slice k-1 goes to the transition table, atemporal heads go to the
257 // atemporal table (guaranteed atemporal-tailed too — a temporal variable can
258 // never be a parent of an atemporal one), everything else (slices 0..k-2,
259 // whether the tail is temporal or atemporal) goes to the initial table.
260 DAG transitionDAG;
261 DAG initialDAG;
262 DAG atemporalDAG;
263
264 const std::size_t nbTransNodes = _transitionLearner_->database().nbVariables();
265 for (std::size_t id = 0; id < nbTransNodes; ++id)
266 transitionDAG.addNodeWithId(NodeId(id));
267
268 const std::size_t nbInitNodes = _initialLearner_->database().nbVariables();
269 for (std::size_t id = 0; id < nbInitNodes; ++id)
270 initialDAG.addNodeWithId(NodeId(id));
271
272 if (_atemporalLearner_) {
273 const std::size_t nbAtemNodes = _atemporalLearner_->database().nbVariables();
274 for (std::size_t id = 0; id < nbAtemNodes; ++id)
275 atemporalDAG.addNodeWithId(NodeId(id));
276 }
277
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);
283
284 if (headSlice == k - 1) {
285 transitionDAG.addArc(_transitionLearner_->idFromName(tailName),
286 _transitionLearner_->idFromName(headName));
287 } else if (headSlice == KTBN< GUM_SCALAR >::ATEMPORAL) {
289 atemporalDAG.addArc(_atemporalLearner_->idFromName(tailName),
290 _atemporalLearner_->idFromName(headName));
291 } else {
292 initialDAG.addArc(_initialLearner_->idFromName(tailName),
293 _initialLearner_->idFromName(headName));
294 }
295 }
296
297 BayesNet< GUM_SCALAR > transitionBN
298 = _transitionLearner_->learnParameters(transitionDAG, takeIntoAccountScore);
299 BayesNet< GUM_SCALAR > initialBN
300 = _initialLearner_->learnParameters(initialDAG, takeIntoAccountScore);
301 BayesNet< GUM_SCALAR > atemporalBN
303 ? _atemporalLearner_->learnParameters(atemporalDAG, takeIntoAccountScore)
304 : BayesNet< GUM_SCALAR >{};
305
306 return _assemble_(transitionBN, initialBN, atemporalBN);
307 }
308
309 // =========================================================================
310 // Score selection
311 // =========================================================================
312
313
314 template < GUM_Numeric GUM_SCALAR >
316 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreAIC(); });
317 return *this;
318 }
319
320 template < GUM_Numeric GUM_SCALAR >
322 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreBD(); });
323 return *this;
324 }
325
326 template < GUM_Numeric GUM_SCALAR >
328 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreBDeu(); });
329 return *this;
330 }
331
332 template < GUM_Numeric GUM_SCALAR >
334 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreBIC(); });
335 return *this;
336 }
337
338 template < GUM_Numeric GUM_SCALAR >
340 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreLog2Likelihood(); });
341 return *this;
342 }
343
344 template < GUM_Numeric GUM_SCALAR >
346 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScoreMDL(); });
347 return *this;
348 }
349
350 template < GUM_Numeric GUM_SCALAR >
352 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useScorefNML(); });
353 }
354
355 template < GUM_Numeric GUM_SCALAR >
357 // Every score/prior setter is applied uniformly to all internal learners, so
358 // their configurations are identical: a single check on the transition learner
359 // is representative of the whole KTBNLearner.
360 return _transitionLearner_->checkScorePriorCompatibility();
361 }
362
363 // =========================================================================
364 // Algorithm selection
365 // =========================================================================
366
367 template < GUM_Numeric GUM_SCALAR >
369 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useGreedyHillClimbing(); });
370 return *this;
371 }
372
373 template < GUM_Numeric GUM_SCALAR >
375 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useExtendedGreedyHillClimbing(); });
376 return *this;
377 }
378
379 template < GUM_Numeric GUM_SCALAR >
383 [&](BNLearner< GUM_SCALAR >& l) { l.useLocalSearchWithTabuList(tabu_size, nb_decrease); });
384 return *this;
385 }
386
387 template < GUM_Numeric GUM_SCALAR >
389 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useMIIC(); });
390 return *this;
391 }
392
393 // =========================================================================
394 // MIIC correction
395 // =========================================================================
396
397 template < GUM_Numeric GUM_SCALAR >
399 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useNMLCorrection(); });
400 return *this;
401 }
402
403 template < GUM_Numeric GUM_SCALAR >
405 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useMDLCorrection(); });
406 return *this;
407 }
408
409 template < GUM_Numeric GUM_SCALAR >
411 _forEachLearner_([](BNLearner< GUM_SCALAR >& l) { l.useNoCorrection(); });
412 return *this;
413 }
414
415 template < GUM_Numeric GUM_SCALAR >
416 std::vector< std::pair< std::string, std::string > >
418 // Each learner reports latent arcs in its own NodeId space, so translate each
419 // Arc to an (engine-name, engine-name) pair before merging. A pair among
420 // slices 0..k-2 can be flagged by both the transition and the initial learner,
421 // so deduplicate on a tab-joined key (engine names never contain a tab).
422 // Propagates BNLearner's OperationNotAllowed if MIIC is not selected.
423 std::vector< std::pair< std::string, std::string > > result;
424 std::unordered_set< std::string > seen;
425
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));
433 }
434 };
435 collect(_transitionLearner_);
436 collect(_initialLearner_);
437 collect(_atemporalLearner_);
438 return result;
439 }
440
441 // =========================================================================
442 // Prior selection
443 // =========================================================================
444
445 template < GUM_Numeric GUM_SCALAR >
447 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.useSmoothingPrior(weight); });
448 return *this;
449 }
450
451 // =========================================================================
452 // Structural constraints
453 // =========================================================================
454
455 template < GUM_Numeric GUM_SCALAR >
457 std::string_view headNode) {
458 _forOwningLearner_(tailNode, headNode, [&](BNLearner< GUM_SCALAR >& l) {
459 l.addForbiddenArc(tailNode, headNode);
460 });
461 return *this;
462 }
463
464 template < GUM_Numeric GUM_SCALAR >
466 int tailSlice,
467 std::string_view headBase,
468 int headSlice) {
469 // wrapper: encode (base, slice) -> engine name and delegate to the string overload
470 return addForbiddenArc(_encode_(tailBase, tailSlice), _encode_(headBase, headSlice));
471 }
472
473 template < GUM_Numeric GUM_SCALAR >
476 std::string_view headNode) {
477 // refuse to lift an invariant the k-TBN definition enforces (shared with addMandatoryArc)
478 _checkArcTemporallyFeasible_(tailNode, headNode, "un-forbid the arc");
479
480 _forOwningLearner_(tailNode, headNode, [&](BNLearner< GUM_SCALAR >& l) {
481 l.eraseForbiddenArc(tailNode, headNode);
482 });
483 return *this;
484 }
485
486 template < GUM_Numeric GUM_SCALAR >
488 int tailSlice,
489 std::string_view headBase,
490 int headSlice) {
491 return eraseForbiddenArc(_encode_(tailBase, tailSlice), _encode_(headBase, headSlice));
492 }
493
494 template < GUM_Numeric GUM_SCALAR >
496 std::string_view headNode) {
497 // Unlike forbidding, a mandatory arc must be feasible: reject up front the arcs
498 // the k-TBN can never contain (shared with eraseForbiddenArc) rather than
499 // letting them crash later in the wrong learner. Kept here, not in
500 // _forOwningLearner_: the check is deliberately asymmetric across the four
501 // setters (see its declaration).
502 _checkArcTemporallyFeasible_(tailNode, headNode, "force the mandatory arc");
503
504 _forOwningLearner_(tailNode, headNode, [&](BNLearner< GUM_SCALAR >& l) {
505 l.addMandatoryArc(tailNode, headNode);
506 });
507 return *this;
508 }
509
510 template < GUM_Numeric GUM_SCALAR >
512 int tailSlice,
513 std::string_view headBase,
514 int headSlice) {
515 // wrapper: encode (base, slice) -> engine name and delegate to the string overload
516 return addMandatoryArc(_encode_(tailBase, tailSlice), _encode_(headBase, headSlice));
517 }
518
519 template < GUM_Numeric GUM_SCALAR >
522 std::string_view headNode) {
523 // no feasibility check: a backward-in-time arc can never have been added, so
524 // erasing it is harmless -- it resolves to a no-op on the owning learner.
525 _forOwningLearner_(tailNode, headNode, [&](BNLearner< GUM_SCALAR >& l) {
526 l.eraseMandatoryArc(tailNode, headNode);
527 });
528 return *this;
529 }
530
531 template < GUM_Numeric GUM_SCALAR >
533 int tailSlice,
534 std::string_view headBase,
535 int headSlice) {
536 return eraseMandatoryArc(_encode_(tailBase, tailSlice), _encode_(headBase, headSlice));
537 }
538
539 template < GUM_Numeric GUM_SCALAR >
542 std::string_view headBase) {
543 // validate both endpoints BEFORE touching any learner, so a rejected call
544 // leaves no partially-applied constraint behind
545 _checkBaseIsTemporal_(tailBase, "an intra-slice constraint");
546 _checkBaseIsTemporal_(headBase, "an intra-slice constraint");
547 const int k = static_cast< int >(_prior_ktbn_.k());
548 for (int t = 0; t < k; ++t)
549 addForbiddenArc(tailBase, t, headBase, t);
550 return *this;
551 }
552
553 template < GUM_Numeric GUM_SCALAR >
556 std::string_view headBase) {
557 _checkBaseIsTemporal_(tailBase, "an intra-slice constraint");
558 _checkBaseIsTemporal_(headBase, "an intra-slice constraint");
559 const int k = static_cast< int >(_prior_ktbn_.k());
560 for (int t = 0; t < k; ++t)
561 eraseForbiddenArc(tailBase, t, headBase, t);
562 return *this;
563 }
564
565 template < GUM_Numeric GUM_SCALAR >
568 std::string_view headBase) {
569 _forEachAllSlicesPair_(tailBase, headBase, [&](int ts, int hs) {
570 addForbiddenArc(tailBase, ts, headBase, hs);
571 });
572 return *this;
573 }
574
575 template < GUM_Numeric GUM_SCALAR >
578 std::string_view headBase) {
579 _forEachAllSlicesPair_(tailBase, headBase, [&](int ts, int hs) {
580 eraseForbiddenArc(tailBase, ts, headBase, hs);
581 });
582 return *this;
583 }
584
585 template < GUM_Numeric GUM_SCALAR >
587 int slice) {
588 return addNoParentNode(_encode_(base, slice));
589 }
590
591 template < GUM_Numeric GUM_SCALAR >
593 const int slice = _determineNode_(std::string{name}).second;
594 const bool isAtemp = (slice == KTBN< GUM_SCALAR >::ATEMPORAL);
595 if (isAtemp && _atemporalLearner_) {
596 // Atemporals are already forced roots in transition/initial learners by _build_.
597 // The user constraint is only meaningful in the atemporal learner.
598 _atemporalLearner_->addNoParentNode(name);
599 } else if (!isAtemp) {
600 _transitionLearner_->addNoParentNode(name);
601 if (slice < (int)_prior_ktbn_.k() - 1) _initialLearner_->addNoParentNode(name);
602 }
603 return *this;
604 }
605
606 template < GUM_Numeric GUM_SCALAR >
608 int slice) {
609 return eraseNoParentNode(_encode_(base, slice));
610 }
611
612 template < GUM_Numeric GUM_SCALAR >
614 const int slice = _determineNode_(std::string{name}).second;
615 const int last = (int)_prior_ktbn_.k() - 1;
616 const bool isAtemp = (slice == KTBN< GUM_SCALAR >::ATEMPORAL);
617 if (isAtemp && _atemporalLearner_) _atemporalLearner_->eraseNoParentNode(name);
618 else if (slice == last) _transitionLearner_->eraseNoParentNode(name);
619 else if (!isAtemp)
620 // Past slices keep their transition-learner root constraint (confines learning
621 // to the kernel); only lift the constraint in the initial learner.
622 _initialLearner_->eraseNoParentNode(name);
623 return *this;
624 }
625
626 template < GUM_Numeric GUM_SCALAR >
628 int slice) {
629 return addNoChildrenNode(_encode_(base, slice));
630 }
631
632 template < GUM_Numeric GUM_SCALAR >
634 const int slice = _determineNode_(std::string{name}).second;
635 const bool isAtemp = (slice == KTBN< GUM_SCALAR >::ATEMPORAL);
636 _transitionLearner_->addNoChildrenNode(name);
637 if (slice < (int)_prior_ktbn_.k() - 1) _initialLearner_->addNoChildrenNode(name);
638 if (isAtemp && _atemporalLearner_) _atemporalLearner_->addNoChildrenNode(name);
639 return *this;
640 }
641
642 template < GUM_Numeric GUM_SCALAR >
644 int slice) {
645 return eraseNoChildrenNode(_encode_(base, slice));
646 }
647
648 template < GUM_Numeric GUM_SCALAR >
650 const int slice = _determineNode_(std::string{name}).second;
651 const bool isAtemp = (slice == KTBN< GUM_SCALAR >::ATEMPORAL);
652 _transitionLearner_->eraseNoChildrenNode(name);
653 if (slice < (int)_prior_ktbn_.k() - 1) _initialLearner_->eraseNoChildrenNode(name);
654 if (isAtemp && _atemporalLearner_) _atemporalLearner_->eraseNoChildrenNode(name);
655 return *this;
656 }
657
658 template < GUM_Numeric GUM_SCALAR >
660 int tailSlice,
661 std::string_view headBase,
662 int headSlice) {
663 // encode (base, slice) -> engine name and delegate to the string overload, which
664 // owns the full routing (atemporal->atemporal to the atemporal learner, slice-k-1 guard)
665 return addPossibleEdge(_encode_(tailBase, tailSlice), _encode_(headBase, headSlice));
666 }
667
668 template < GUM_Numeric GUM_SCALAR >
670 int tailSlice,
671 std::string_view headBase,
672 int headSlice) {
673 return erasePossibleEdge(_encode_(tailBase, tailSlice), _encode_(headBase, headSlice));
674 }
675
676 template < GUM_Numeric GUM_SCALAR >
678 std::string_view head) {
679 const int last = (int)_prior_ktbn_.k() - 1;
680 const int tailSlice = _determineNode_(std::string{tail}).second;
681 const int headSlice = _determineNode_(std::string{head}).second;
682
683 // Atemporal names exist as forced roots in every table, so whitelisting them
684 // here too keeps the whole-model promise ("only listed edges are explored")
685 // even when the edge can fire in just one learner: the no-parent constraint
686 // on atemporal heads still vetoes it there (constraints compound, never
687 // override), so this can only shrink the candidate set, never widen it.
688 _transitionLearner_->addPossibleEdge(tail, head);
689 // forward to the initial learner only when both endpoints exist there, i.e. both
690 // are at a past slice (< k-1); atemporals (slice -1) satisfy this automatically.
691 if (tailSlice < last && headSlice < last) _initialLearner_->addPossibleEdge(tail, head);
692
693 if (tailSlice == KTBN< GUM_SCALAR >::ATEMPORAL && headSlice == KTBN< GUM_SCALAR >::ATEMPORAL) {
695 if (_atemporalLearner_) _atemporalLearner_->addPossibleEdge(tail, head);
696 } else {
698 }
699 return *this;
700 }
701
702 template < GUM_Numeric GUM_SCALAR >
704 std::string_view head) {
705 const int last = (int)_prior_ktbn_.k() - 1;
706 const int tailSlice = _determineNode_(std::string{tail}).second;
707 const int headSlice = _determineNode_(std::string{head}).second;
708
709 _transitionLearner_->erasePossibleEdge(tail, head);
710 if (tailSlice < last && headSlice < last) _initialLearner_->erasePossibleEdge(tail, head);
711
712 if (tailSlice == KTBN< GUM_SCALAR >::ATEMPORAL && headSlice == KTBN< GUM_SCALAR >::ATEMPORAL) {
714 if (_atemporalLearner_) _atemporalLearner_->erasePossibleEdge(tail, head);
715 } else {
717 }
718 return *this;
719 }
720
721 template < GUM_Numeric GUM_SCALAR >
723 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.allowArcAdditions(allow); });
724 return *this;
725 }
726
727 template < GUM_Numeric GUM_SCALAR >
729 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.allowArcDeletions(allow); });
730 return *this;
731 }
732
733 template < GUM_Numeric GUM_SCALAR >
735 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.allowArcReversals(allow); });
736 return *this;
737 }
738
739 template < GUM_Numeric GUM_SCALAR >
741 _forEachLearner_([&](BNLearner< GUM_SCALAR >& l) { l.setMaxIndegree(max_indegree); });
742 return *this;
743 }
744
745 // =========================================================================
746 // Diagnostics
747 // =========================================================================
748
749 template < GUM_Numeric GUM_SCALAR >
751 return _prior_ktbn_.k();
752 }
753
754 template < GUM_Numeric GUM_SCALAR >
756 return _prior_ktbn_.nbTemporalVars() + _prior_ktbn_.nbAtemporalVars();
757 }
758
759 template < GUM_Numeric GUM_SCALAR >
760 std::vector< Size > KTBNLearner< GUM_SCALAR >::nbRows() const {
761 // Raw trajectory lengths captured by _build_(), one entry per sample. Unlike
762 // BNLearner::nbRows() (a single flat-table row count), a trajectory learner
763 // has one length per sequence, so the per-sample vector is the natural analog.
764 return _nbTimeSlices_;
765 }
766
767 template < GUM_Numeric GUM_SCALAR >
769 return _transitionLearner_->isConstraintBased();
770 }
771
772 template < GUM_Numeric GUM_SCALAR >
774 return _transitionLearner_->isScoreBased();
775 }
776
777 template < GUM_Numeric GUM_SCALAR >
779 // Emit a k-TBN-specific header, then delegate score/algo/prior/constraint
780 // details to each internal learner's toString().
781 std::stringstream s;
782 s << "k : " << k() << '\n';
783 s << "Variables : " << nbCols() << " (" << _prior_ktbn_.nbTemporalVars()
784 << " temporal, " << _prior_ktbn_.nbAtemporalVars() << " atemporal)" << '\n';
785 s << "Transition rows : " << _transitionLearner_->nbRows() << '\n';
786 s << "Initial rows : " << _initialLearner_->nbRows() << '\n';
787 s << '\n';
788 s << "=== Transition learner (arcs into slice " << (k() - 1) << ") ===" << '\n';
789 s << _transitionLearner_->toString();
790 s << '\n';
791 s << "=== Initial learner (slices 0.." << (k() - 2) << ") ===" << '\n';
792 s << _initialLearner_->toString();
793 if (_atemporalLearner_) {
794 s << '\n';
795 s << "=== Atemporal learner (atemporal->atemporal arcs) ===" << '\n';
796 s << _atemporalLearner_->toString();
797 }
798 return s.str();
799 }
800
801 template < GUM_Numeric GUM_SCALAR >
802 std::vector< std::tuple< std::string, std::string, std::string > >
804 // settings are applied identically to every internal learner, so the
805 // transition learner's state represents the whole KTBNLearner.
806 auto result = _transitionLearner_->state();
807
808 // The transition learner lists every slice of every temporal variable
809 // (e.g. "W[0][2], W[1][2], W[2][2], W[3][2]"). Replace that row with one
810 // entry per base variable (e.g. "W[2]") built directly from the KTBN's own
811 // variable sets — no string parsing, robust to any base name.
812 for (auto& [key, val, comment]: result) {
813 if (key != "Variables") continue;
814 // Build base -> domainSize from the prior KTBN (slice 0 for temporal, AT for atemporal).
815 std::string collapsed;
816 bool first = true;
817 auto emit = [&](const std::string& base, int slice) {
818 const auto& var = _prior_ktbn_.variable(base, slice);
819 if (!first) collapsed += ", ";
820 collapsed += base + "[" + std::to_string(var.domainSize()) + "]";
821 first = false;
822 };
823 for (const auto& base: _prior_ktbn_.temporalVarNames())
824 emit(base, 0);
825 for (const auto& base: _prior_ktbn_.atemporalVarNames())
827 val = collapsed;
828 break;
829 }
830 return result;
831 }
832
833 template < GUM_Numeric GUM_SCALAR >
835 // As in BNLearner::copyState, only score/algorithm/prior/constraint settings
836 // are copied, never the database. Both sides must be structurally compatible
837 // (same k, same temporal/atemporal partition) for the name-based constraints
838 // to mean the same thing. The atemporal learner is copied only when both
839 // sides have one (absent when there are <= 1 atemporal vars).
840 _transitionLearner_->copyState(*learner._transitionLearner_);
841 _initialLearner_->copyState(*learner._initialLearner_);
843 _atemporalLearner_->copyState(*learner._atemporalLearner_);
844
845 // KTBNLearner-only state (BNLearner::copyState can't carry it): learnKTBN reads
846 // these to suppress the atemporal learner under a temporal-only whitelist
849 }
850
851 // =========================================================================
852 // Database accessors
853 // =========================================================================
854
855 template < GUM_Numeric GUM_SCALAR >
857 // the initial table holds exactly one row per trajectory, so its row count
858 // is the number of trajectory files given to the constructor.
859 return _initialLearner_->nbRows();
860 }
861
862 template < GUM_Numeric GUM_SCALAR >
866
867 template < GUM_Numeric GUM_SCALAR >
871
872 template < GUM_Numeric GUM_SCALAR >
874 return _transitionLearner_->hasMissingValues() || _initialLearner_->hasMissingValues()
875 || (_atemporalLearner_ && _atemporalLearner_->hasMissingValues());
876 }
877
878 template < GUM_Numeric GUM_SCALAR >
879 std::vector< std::string > KTBNLearner< GUM_SCALAR >::names() const {
880 // The first nbCols() columns of the transition table are the base variables
881 // (slice-0 temporals then atemporals); the remaining columns are just later
882 // slices of those same temporals. So strip the slice suffix off the first
883 // nbCols() engine names to get each base variable exactly once.
884 const auto& engine = _transitionLearner_->names();
885 const Size n = nbCols();
886 std::vector< std::string > result;
887 result.reserve(n);
888 for (Size i = 0; i < n; ++i)
889 result.push_back(_determineNode_(engine[i]).first);
890 return result;
891 }
892
893 template < GUM_Numeric GUM_SCALAR >
894 std::vector< Size > KTBNLearner< GUM_SCALAR >::domainSizes() const {
895 // Mirror names(): the first nbCols() columns are the base variables (a temporal
896 // variable's domain size is the same across all its slices), so keep only those.
897 const auto& all = _transitionLearner_->domainSizes();
898 return std::vector< Size >(all.begin(), all.begin() + nbCols());
899 }
900
901 template < GUM_Numeric GUM_SCALAR >
902 Size KTBNLearner< GUM_SCALAR >::domainSize(std::string_view base) const {
903 // Resolve the base name to an engine name in the transition table: a temporal
904 // variable is addressed at slice 0, everything else keeps its bare name (which
905 // also lets a raw engine name pass through). An unknown name reaches
906 // domainSize() unchanged and throws MissingVariableInDatabase.
907 const std::string b{base};
908 const int slice
909 = _prior_ktbn_.temporalVarNames().contains(b) ? 0 : KTBN< GUM_SCALAR >::ATEMPORAL;
910 return _transitionLearner_->domainSize(_encode_(b, slice));
911 }
912
913 // =========================================================================
914 // Private helpers
915 // =========================================================================
916
917
918 template < GUM_Numeric GUM_SCALAR >
919 std::unordered_set< std::string > KTBNLearner< GUM_SCALAR >::_inferAtemporalVars_(
920 std::string_view dirPath,
921 std::string_view csvBaseName,
923 Size k,
924 const std::vector< std::string >& missingSymbols) {
927 csvBaseName,
928 nbSamples,
929 missingSymbols);
930 }
931
932 template < GUM_Numeric GUM_SCALAR >
934 std::string_view dirPath,
935 std::string_view csvBaseName,
936 Size k,
937 const std::unordered_set< std::string >& atemporalVars,
938 const std::vector< std::string >& missingSymbols,
939 bool induceTypes) {
941
942 namespace fs = std::filesystem;
943 const std::string firstCSV
944 = (fs::path{dirPath} / (std::string{csvBaseName} + "1.csv")).string();
945 // _build_() also reads trajectory 1 (its i=0 pass). The double-read is
946 // unavoidable: this function runs in the member-initialiser list, before
947 // the object (and hence _build_()) exists, and _build_() must read every
948 // trajectory.
949
950 // BNLearner construction runs induceTypes on the CSV and populates all
951 // column translators with properly-typed DiscreteVariables.
952 const BNLearner< GUM_SCALAR > tmpLearner(firstCSV, missingSymbols, induceTypes);
953
954 // Pull typed variables from the translator set. translatorSafe(i) bounds-checks,
955 // guarding against any names()/translator-set mismatch.
956 const DBTranslatorSet& translators = tmpLearner.database().translatorSet();
957 const std::vector< std::string >& names = tmpLearner.names();
958
959 KTBN< GUM_SCALAR > prior(k);
960 for (std::size_t i = 0; i < names.size(); ++i) {
961 prior.add(static_cast< const DiscreteVariable& >(*translators.translatorSafe(i).variable()),
962 !atemporalVars.contains(names[i]));
963 }
964
965 // Every declared atemporal name must appear in the CSV header.
966 for (const std::string& aname: atemporalVars)
967 if (!prior.exists(aname))
969 "atemporal variable '" << aname << "' not found in the CSV header")
970 return prior;
971 }
972
973 template < GUM_Numeric GUM_SCALAR >
975 Size k,
976 const BayesNet< GUM_SCALAR >& bn,
977 const std::unordered_set< std::string >& atemporalVars) {
979
980 KTBN< GUM_SCALAR > prior(k);
981 // bn.nodes() iterates in hash order (unspecified). This is harmless because
982 // KTBN looks up variables by name, not by insertion index.
983 for (const NodeId node: bn.nodes()) {
984 const DiscreteVariable& var = bn.variable(node);
985 prior.add(var, !atemporalVars.contains(var.name()));
986 }
987 for (const std::string& aname: atemporalVars)
988 if (!prior.exists(aname))
989 GUM_ERROR(InvalidArgument, "atemporal variable '" << aname << "' not found in the BN")
990 return prior;
991 }
992
993 template < GUM_Numeric GUM_SCALAR >
994 void KTBNLearner< GUM_SCALAR >::_build_(std::string_view dirPath,
995 std::string_view csvBaseName,
997 const std::vector< std::string >& missingSymbols) {
998 Size k = _prior_ktbn_.k();
999 Size nbTempVars = _prior_ktbn_.temporalVarNames().size();
1000 Size nbAtempVars = _prior_ktbn_.atemporalVarNames().size();
1001
1002 // tables are default-constructed here; translators are inserted per-column
1003 // in the i==0 block once the header (column order) is known
1004 DatabaseTable transitionTable(missingSymbols);
1005 DatabaseTable initTable(missingSymbols);
1006 DatabaseTable atemporalTable(missingSymbols);
1007
1008 // Complete-case selection, at row granularity.
1009 //
1010 // aGrUM's structure learning refuses a database holding any missing value
1011 // outright (IBNLearner::learnDag_), so an incomplete row cannot simply be
1012 // handed over: learnKTBN() would throw on any trajectory with a gap. A row
1013 // that carries one is therefore dropped here, and everything downstream sees
1014 // a database with no missing value at all.
1015 //
1016 // This mirrors what the cross-k score already does per instance
1017 // (KTBNAdaptiveLearner::_forEachScoredNode_ skips any instance whose family
1018 // is not fully observed), so the two layers agree on which data counts.
1019 //
1020 // The unit differs per table: a transition row is a whole width-k window, so
1021 // one gap costs up to k windows; an initial or atemporal row is the whole
1022 // trajectory's contribution to that table.
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)) {
1030 return;
1031 }
1032 }
1033 table.insertRow(r);
1034 };
1035
1036 const std::filesystem::path dir{dirPath};
1037 const std::string stem{csvBaseName};
1038
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;
1042
1043 std::vector< std::string > header; // column order, captured from trajectory 1
1044 std::unordered_set< Size > atemVarsCols; // atemporal column indices
1045
1046 // reused across trajectories (capacity is kept between iterations)
1047 std::vector< std::vector< std::string > > buffer;
1048 std::vector< std::string > row;
1049 row.reserve(transRowSize);
1050 _nbTimeSlices_.reserve(nbSamples); // one raw length per trajectory (see nbRows())
1051
1052 for (Size i = 0; i < nbSamples; ++i) {
1053 // open the i-th trajectory (each file is opened exactly once)
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());
1057
1058 CSVParser parser(is, file.string());
1059 parser.next();
1060
1061 if (i == 0) {
1062 // first file only: capture the column order, derive the atemporal columns
1063 // from _prior_ktbn_, build the bracket-encoded names and finish setting up the
1064 // tables (must happen before any insertRow)
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);
1069
1070 // every schema variable must appear in the data, else _assemble_ (used
1071 // by both learnKTBN and learnParameters) would later fail with an
1072 // opaque NotFound when resolving bracket names
1073 {
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() << "'")
1079 };
1080 for (const auto& base: _prior_ktbn_.temporalVarNames())
1081 requirePresent(base);
1082 for (const auto& base: _prior_ktbn_.atemporalVarNames())
1083 requirePresent(base);
1084
1085 // conversely, every CSV column must be a known schema variable, else the
1086 // translator-insertion loop below would fail with an opaque NotFound
1087 // when resolving it against _prior_ktbn_
1088 for (const std::string& col: header)
1089 if (!_prior_ktbn_.exists(col))
1091 "CSV column '" << col << "' in '" << file.string()
1092 << "' is not declared as a variable of this KTBNLearner")
1093 }
1094
1095 // insert one translator per column into each table; every schema variable
1096 // has a concrete domain (user-supplied, or discovered by _buildPriorFromCSV_), so
1097 // insertTranslator picks the matching translator type from the variable.
1098 std::vector< std::string > varNamesTran;
1099 std::vector< std::string > varNamesInit;
1100 std::vector< std::string > varNamesAtemp;
1101 {
1102 auto insertTrans
1103 = [&](DatabaseTable& table, const std::string& base, int slice, std::size_t col) {
1104 table.insertTranslator(_prior_ktbn_.variable(base, slice), col, missingSymbols);
1105 };
1106
1107 // translators and their engine names are built in lockstep, column by
1108 // column, so a translator's table position and its name can never drift
1109 // apart (unlike keeping two separately-indexed passes in sync by hand).
1110 varNamesTran.reserve(transRowSize);
1111 varNamesInit.reserve(initRowSize);
1112 varNamesAtemp.reserve(atempRowSize);
1113
1114 std::size_t tcol = 0, icol = 0, acol = 0;
1115 for (std::size_t c = 0; c < header.size(); ++c) {
1116 const int slice = atemVarsCols.contains(c) ? KTBN< GUM_SCALAR >::ATEMPORAL : 0;
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)) {
1122 insertTrans(atemporalTable, header[c], KTBN< GUM_SCALAR >::ATEMPORAL, acol++);
1123 varNamesAtemp.push_back(header[c]);
1124 }
1125 }
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));
1131 }
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));
1137 }
1138 }
1139
1140 transitionTable.setVariableNames(varNamesTran, false);
1141 initTable.setVariableNames(varNamesInit, false);
1142 atemporalTable.setVariableNames(varNamesAtemp, false);
1143 } else {
1144 // later files: validate the header against trajectory 1's without copying it
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]);
1149 if (!same)
1150 GUM_ERROR(InvalidArgument, "Header of " << file.string() << " differs from trajectory 1");
1151 }
1152
1153 buffer.clear();
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()});
1161 }
1162
1163 if (buffer.size() < k)
1165 "Trajectory " << (i + 1) << " has " << buffer.size()
1166 << " time steps but at least k=" << k << " are required");
1167
1168 // record this trajectory's raw length (number of time steps), exposed by nbRows()
1169 _nbTimeSlices_.push_back(buffer.size());
1170
1171 // An atemporal column is constant down the trajectory, so any row carries
1172 // its value -- but row 0's may be the missing one, and reading row 0 blindly
1173 // would then drop every row this trajectory feeds. Resolve each once, from
1174 // its first non-missing occurrence. Same rule the score walk applies.
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];
1180 break;
1181 }
1182
1183 // value of column col at time tt: the resolved constant for an atemporal
1184 // column, the row's own cell for a temporal one. A column missing all the
1185 // way down has no resolved value, so row 0's marker stands and the row is
1186 // dropped like any other incomplete one.
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;
1191 };
1192
1193 // transition table: sliding windows of width k
1194 for (std::size_t t = 0; t + k <= buffer.size(); ++t) {
1195 row.clear();
1196 // slice 0 carries the atemporal columns; later slices skip them so atemporal
1197 // variables aren't repeated k-1 times in each row
1198 for (Size col = 0; col < buffer[0].size(); ++col) {
1199 row.push_back(cellAt(t, col));
1200 }
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]); }
1204 }
1205 }
1206 insertIfComplete(transitionTable, row);
1207 }
1208
1209 // initial table: first k-1 slices, one row per trajectory
1210 row.clear();
1211 for (Size col = 0; col < buffer[0].size(); ++col) {
1212 row.push_back(cellAt(0, col));
1213 }
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]); }
1217 }
1218 }
1219 insertIfComplete(initTable, row);
1220
1221 // atemporal table: one row per trajectory, atemporal columns only. Their
1222 // value is constant across the trajectory, so any time step works — take 0.
1223 // Skipped when fewer than 2 atemporal variables exist: a single atemporal
1224 // variable has no possible atemporal->atemporal arcs, so no learner is built.
1225 if (nbAtempVars > 1) {
1226 row.clear();
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);
1230 }
1231 }
1232
1233 // Dropping incomplete rows can empty a table outright -- every transition
1234 // window straddling a gap, or every trajectory's initial block incomplete.
1235 // The internal learner would then fail obscurely on an empty database, so
1236 // say what actually happened.
1237 if (_ignoreMissingSymbols_ && (transitionTable.nbRows() == 0 || initTable.nbRows() == 0))
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="
1243 << k << ".")
1244
1245 // Variables are already typed upstream (template / first-trajectory learner),
1246 // so no induceTypes pass is needed here — just canonicalize the value codes.
1247 transitionTable.reorder();
1248 initTable.reorder();
1249
1250 _transitionLearner_ = std::make_unique< BNLearner< GUM_SCALAR > >(transitionTable);
1251 _initialLearner_ = std::make_unique< BNLearner< GUM_SCALAR > >(initTable);
1252
1253 if (nbAtempVars > 1) {
1254 atemporalTable.reorder();
1255 _atemporalLearner_ = std::make_unique< BNLearner< GUM_SCALAR > >(atemporalTable);
1256 }
1257
1258
1259 // Impose k-TBN temporal constraints on structure learning:
1260 // - transition learner: past slices (0..k-2) are forced roots so only the
1261 // present slice (k-1) receives new arcs.
1262 // - atemporal variables are forced roots in transition/initial learners so
1263 // their mutual structure is learned exclusively by the atemporal learner.
1264 // No-parent in both learners covers every algorithm and also bans
1265 // temporal→atemporal arcs.
1266 // - initial learner: backward temporal arcs among past slices forbidden
1267 // explicitly (honoured by MIIC and score-based algorithms alike).
1268 // These constraints shape structure search in learnKTBN(); they are inert
1269 // for learnParameters.
1270 const int ki = (int)k;
1271 const auto& temporalVars = _prior_ktbn_.temporalVarNames();
1272 const auto& atemporalVars = _prior_ktbn_.atemporalVarNames();
1273
1274 for (const auto& base: temporalVars)
1275 for (int slice = 0; slice < ki - 1; ++slice)
1276 _transitionLearner_->addNoParentNode(_encode_(base, slice));
1277
1278 for (const auto& atemBase: atemporalVars) {
1279 _transitionLearner_->addNoParentNode(atemBase);
1280 _initialLearner_->addNoParentNode(atemBase);
1281 }
1282
1283 if (k > 2) {
1284 // backward-in-time arcs among the past slices 0..k-2 are forbidden
1285 // explicitly: this is the only form MIIC honours (it ignores slice order),
1286 // and score-based algorithms respect it too, so it fully covers the
1287 // constraint. A setSliceOrder() mirror was dropped here as redundant.
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)
1292 _initialLearner_->addForbiddenArc(_encode_(tailBase, tailSlice),
1293 _encode_(headBase, headSlice));
1294 }
1295 }
1296
1297 template < GUM_Numeric GUM_SCALAR >
1298 const std::unordered_set< std::string >& KTBNLearner< GUM_SCALAR >::_atemporalVarNames_() const {
1299 return _prior_ktbn_.atemporalVarNames();
1300 }
1301
1302 template < GUM_Numeric GUM_SCALAR >
1303 bool KTBNLearner< GUM_SCALAR >::_isKnownBase_(std::string_view base) const {
1304 const std::string b{base};
1305 return _prior_ktbn_.temporalVarNames().contains(b)
1306 || _prior_ktbn_.atemporalVarNames().contains(b);
1307 }
1308
1309 template < GUM_Numeric GUM_SCALAR >
1310 KTBN< GUM_SCALAR >
1311 KTBNLearner< GUM_SCALAR >::_assemble_(const BayesNet< GUM_SCALAR >& transitionBN,
1312 const BayesNet< GUM_SCALAR >& initialBN,
1313 const BayesNet< GUM_SCALAR >& atemporalBN) const {
1314 const int k = (int)_prior_ktbn_.k();
1315
1316 // 1. Base: initialBN provides slices 0..k-2 (temporal structure + CPTs, and
1317 // atemporal→temporal arcs). Atemporal variables enter only as forced roots;
1318 // steps 4 and 6 overwrite their mutual structure and CPTs from atemporalBN
1319 // (transitionBN repeats each atemporal value k-fold — not independent samples).
1320 BayesNet< GUM_SCALAR > bn(initialBN);
1321
1322 // 2. Add the slice k-1 variables from transitionBN (absent from 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));
1327 }
1328
1329 // 3. Add the arcs arriving at slice k-1 from transitionBN
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);
1336 }
1337 }
1338
1339 // 4. Add the atemporal -> atemporal arcs from atemporalBN. Skipped when
1340 // atemporalBN is an empty placeholder: either <= 1 atemporal variable (no
1341 // learner is built, see _build_()) or the learner was suppressed by a
1342 // temporal-only possible-edge whitelist (see learnKTBN()). Either way the
1343 // atemporal variables already sit in bn as roots, via initialBN.
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);
1350 }
1351 }
1352 }
1353
1354 // 5. Fill the slice k-1 CPTs from transitionBN
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));
1360 }
1361
1362 // 6. Fill the atemporal CPTs from atemporalBN (same guard as step 4); when
1363 // skipped, the atemporal variables keep their initialBN marginal.
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));
1369 }
1370 }
1371
1372 // 7. Convert the assembled flat BN into a KTBN: fromBN infers k from the
1373 // highest bracket index and validates temporal causality.
1374 return KTBN< GUM_SCALAR >::fromBN(bn);
1375 }
1376
1377} // namespace gum::learning
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.
Base class for dag.
Definition DAG.h:121
void addArc(NodeId tail, NodeId head) final
insert a new arc into the directed graph
Definition DAG_inl.h:75
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.
Definition KTBN.h:200
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
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).
Definition CSVParser.h:78
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.
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)
Definition exceptions.h:76
std::size_t Size
In aGrUM, hashed values are unsigned long int.
Definition types.h:74
Size NodeId
Type for node ids.
include the inlined functions if necessary
Definition CSVParser.h:55