50#ifndef DOXYGEN_SHOULD_SKIP_THIS
63 const std::vector< std::pair< std::size_t, std::size_t > >& ranges,
64 const Bijection< NodeId, std::size_t >& nodeId2columns) :
65 _nodeId2columns_(nodeId2columns) {
67 const std::size_t db_nb_cols = parser.database().nbVariables();
68 for (
auto iter = nodeId2columns.cbegin(); iter != nodeId2columns.cend(); ++iter) {
69 if (iter.second() >= db_nb_cols) {
70 GUM_ERROR(OutOfBounds,
71 "the mapping between ids and database columns "
72 <<
"is incorrect because Column " << iter.second()
73 <<
" does not belong to the database.");
78 const auto max_nb_threads = ThreadNumberManager::getNumberOfThreads();
79 _parsers_.reserve(max_nb_threads);
80 for (std::size_t i = std::size_t(0); i < max_nb_threads; ++i)
81 _parsers_.push_back(parser);
85 _checkRanges_(ranges);
86 _ranges_.reserve(ranges.size());
87 for (
const auto& range: ranges)
88 _ranges_.push_back(range);
91 _dispatchRangesToThreads_();
93 GUM_CONSTRUCTOR(RecordCounter);
97 RecordCounter::RecordCounter(
const DBRowGeneratorParser& parser,
98 const Bijection< NodeId, std::size_t >& nodeId2columns) :
104 RecordCounter::RecordCounter(
const RecordCounter& from) :
105 ThreadNumberManager(from), _parsers_(from._parsers_), _ranges_(from._ranges_),
106 _thread_ranges_(from._thread_ranges_), _nodeId2columns_(from._nodeId2columns_),
107 _last_DB_counting_(from._last_DB_counting_), _last_DB_ids_(from._last_DB_ids_),
108 _last_nonDB_counting_(from._last_nonDB_counting_), _last_nonDB_ids_(from._last_nonDB_ids_),
109 _min_nb_rows_per_thread_(from._min_nb_rows_per_thread_) {
110 GUM_CONS_CPY(RecordCounter);
114 RecordCounter::RecordCounter(RecordCounter&& from) :
115 ThreadNumberManager(
std::move(from)), _parsers_(
std::move(from._parsers_)),
116 _ranges_(
std::move(from._ranges_)), _thread_ranges_(
std::move(from._thread_ranges_)),
117 _nodeId2columns_(
std::move(from._nodeId2columns_)),
118 _last_DB_counting_(
std::move(from._last_DB_counting_)),
119 _last_DB_ids_(
std::move(from._last_DB_ids_)),
120 _last_nonDB_counting_(
std::move(from._last_nonDB_counting_)),
121 _last_nonDB_ids_(
std::move(from._last_nonDB_ids_)),
122 _min_nb_rows_per_thread_(from._min_nb_rows_per_thread_) {
123 GUM_CONS_MOV(RecordCounter);
127 RecordCounter* RecordCounter::clone()
const {
return new RecordCounter(*
this); }
130 RecordCounter::~RecordCounter() { GUM_DESTRUCTOR(RecordCounter); }
133 RecordCounter& RecordCounter::operator=(
const RecordCounter& from) {
135 ThreadNumberManager::operator=(from);
136 _parsers_ = from._parsers_;
137 _ranges_ = from._ranges_;
138 _thread_ranges_ = from._thread_ranges_;
139 _nodeId2columns_ = from._nodeId2columns_;
140 _last_DB_counting_ = from._last_DB_counting_;
141 _last_DB_ids_ = from._last_DB_ids_;
142 _last_nonDB_counting_ = from._last_nonDB_counting_;
143 _last_nonDB_ids_ = from._last_nonDB_ids_;
144 _min_nb_rows_per_thread_ = from._min_nb_rows_per_thread_;
150 RecordCounter& RecordCounter::operator=(RecordCounter&& from) {
152 ThreadNumberManager::operator=(std::move(from));
153 _parsers_ = std::move(from._parsers_);
154 _ranges_ = std::move(from._ranges_);
155 _thread_ranges_ = std::move(from._thread_ranges_);
156 _nodeId2columns_ = std::move(from._nodeId2columns_);
157 _last_DB_counting_ = std::move(from._last_DB_counting_);
158 _last_DB_ids_ = std::move(from._last_DB_ids_);
159 _last_nonDB_counting_ = std::move(from._last_nonDB_counting_);
160 _last_nonDB_ids_ = std::move(from._last_nonDB_ids_);
161 _min_nb_rows_per_thread_ = from._min_nb_rows_per_thread_;
167 void RecordCounter::clear() {
168 _last_DB_counting_.clear();
169 _last_DB_ids_.clear();
170 _last_nonDB_counting_.clear();
171 _last_nonDB_ids_.clear();
176 void RecordCounter::setMinNbRowsPerThread(
const std::size_t nb)
const {
177 if (nb == std::size_t(0)) _min_nb_rows_per_thread_ = std::size_t(1);
178 else _min_nb_rows_per_thread_ = nb;
182 void RecordCounter::_raiseCheckException_(
const std::vector< std::string >& bad_vars)
const {
184 if (bad_vars.size() == 1) {
186 std::format(
"Counts cannot be performed on continuous variables. "
187 "Unfortunately the following variable is continuous: {}",
192 for (
const auto& name: bad_vars) {
193 if (!first) varList +=
", ";
198 std::format(
"Counts cannot be performed on continuous variables. "
199 "Unfortunately the following variables are continuous: {}",
205 void RecordCounter::_checkDiscreteVariables_(
const IdCondSet& ids)
const {
206 const std::size_t size = ids.size();
207 const DatabaseTable& database = _parsers_[0].data.database();
209 if (_nodeId2columns_.empty()) {
211 for (std::size_t i = std::size_t(0); i < size; ++i) {
212 if (database.variable(i).varType() == VarType::CONTINUOUS) {
216 std::vector< std::string > bad_vars{database.variable(i).name()};
217 for (++i; i < size; ++i) {
218 if (database.variable(i).varType() == VarType::CONTINUOUS)
219 bad_vars.push_back(database.variable(i).name());
221 _raiseCheckException_(bad_vars);
226 for (std::size_t i = std::size_t(0); i < size; ++i) {
228 std::size_t pos = _nodeId2columns_.second(ids[i]);
230 if (database.variable(pos).varType() == VarType::CONTINUOUS) {
234 std::vector< std::string > bad_vars{database.variable(pos).name()};
235 for (++i; i < size; ++i) {
236 pos = _nodeId2columns_.second(ids[i]);
237 if (database.variable(pos).varType() == VarType::CONTINUOUS)
238 bad_vars.push_back(database.variable(pos).name());
240 _raiseCheckException_(bad_vars);
248 HashTable< NodeId, std::size_t >
249 RecordCounter::_getNodeIds2Columns_(
const IdCondSet& ids)
const {
250 HashTable< NodeId, std::size_t > res(ids.size());
251 if (_nodeId2columns_.empty()) {
252 for (
const auto id: ids) {
253 res.insert(
id, std::size_t(
id));
256 for (
const auto id: ids) {
257 res.insert(
id, _nodeId2columns_.second(
id));
264 std::vector< double >&
265 RecordCounter::_extractFromCountings_(
const IdCondSet& subset_ids,
266 const IdCondSet& superset_ids,
267 const std::vector< double >& superset_vect) {
271 const auto nodeId2columns = _getNodeIds2Columns_(superset_ids);
275 const auto& database = _parsers_[0].data.database();
276 std::size_t result_vect_size = std::size_t(1);
277 for (
const auto id: subset_ids) {
278 result_vect_size *= database.domainSize(nodeId2columns[
id]);
282 std::vector< double > result_vect(result_vect_size, 0.0);
288 bool subset_begin =
true;
289 const std::size_t subset_ids_size = std::size_t(subset_ids.size());
290 for (std::size_t i = 0; i < subset_ids_size; ++i) {
291 if (superset_ids.pos(subset_ids[i]) != i) {
292 subset_begin =
false;
298 const std::size_t superset_vect_size = superset_vect.size();
299 std::size_t i = std::size_t(0);
300 while (i < superset_vect_size) {
301 for (std::size_t j = std::size_t(0); j < result_vect_size; ++j, ++i) {
302 result_vect[j] += superset_vect[i];
308 _last_nonDB_ids_ = subset_ids;
309 _last_nonDB_counting_ = std::move(result_vect);
310 return _last_nonDB_counting_;
312 _last_nonDB_ids_.clear();
313 _last_nonDB_counting_.clear();
322 bool subset_end =
true;
323 const std::size_t superset_ids_size = std::size_t(superset_ids.size());
324 for (std::size_t i = 0; i < subset_ids_size; ++i) {
325 if (superset_ids.pos(subset_ids[i]) != i + superset_ids_size - subset_ids_size) {
334 std::size_t vect_not_subset_size = std::size_t(1);
335 for (std::size_t i = std::size_t(0); i < superset_ids_size - subset_ids_size; ++i)
336 vect_not_subset_size *= database.domainSize(nodeId2columns[superset_ids[i]]);
339 std::size_t i = std::size_t(0);
340 for (std::size_t j = std::size_t(0); j < result_vect_size; ++j) {
341 for (std::size_t k = std::size_t(0); k < vect_not_subset_size; ++k, ++i) {
342 result_vect[j] += superset_vect[i];
348 _last_nonDB_ids_ = subset_ids;
349 _last_nonDB_counting_ = std::move(result_vect);
350 return _last_nonDB_counting_;
352 _last_nonDB_ids_.clear();
353 _last_nonDB_counting_.clear();
384 std::vector< std::size_t > before_incr(subset_ids_size);
385 std::vector< std::size_t > result_domain(subset_ids_size);
386 std::vector< std::size_t > result_offset(subset_ids_size);
388 std::size_t result_domain_size = std::size_t(1);
389 std::size_t tmp_before_incr = std::size_t(1);
390 std::vector< std::size_t > superset_order(subset_ids_size);
392 for (std::size_t h = std::size_t(0), j = std::size_t(0); j < subset_ids_size; ++h) {
393 if (subset_ids.exists(superset_ids[h])) {
394 before_incr[j] = tmp_before_incr - 1;
395 superset_order[subset_ids.pos(superset_ids[h])] = j;
399 tmp_before_incr *= database.domainSize(nodeId2columns[superset_ids[h]]);
404 for (std::size_t i = 0; i < subset_ids.size(); ++i) {
405 const std::size_t domain_size = database.domainSize(nodeId2columns[subset_ids[i]]);
406 const std::size_t j = superset_order[i];
407 result_domain[j] = domain_size;
408 result_offset[j] = result_domain_size;
409 result_domain_size *= domain_size;
413 std::vector< std::size_t > result_value(result_domain);
414 std::vector< std::size_t > current_incr(before_incr);
415 std::vector< std::size_t > result_down(result_offset);
417 for (std::size_t j = std::size_t(0); j < result_down.size(); ++j) {
418 result_down[j] *= (result_domain[j] - 1);
422 const std::size_t superset_vect_size = superset_vect.size();
423 std::size_t the_result_offset = std::size_t(0);
424 for (std::size_t h = std::size_t(0); h < superset_vect_size; ++h) {
425 result_vect[the_result_offset] += superset_vect[h];
428 for (std::size_t k = 0; k < current_incr.size(); ++k) {
430 if (current_incr[k]) {
435 current_incr[k] = before_incr[k];
440 if (result_value[k]) {
441 the_result_offset += result_offset[k];
445 result_value[k] = result_domain[k];
446 the_result_offset -= result_down[k];
452 _last_nonDB_ids_ = subset_ids;
453 _last_nonDB_counting_ = std::move(result_vect);
454 return _last_nonDB_counting_;
456 _last_nonDB_ids_.clear();
457 _last_nonDB_counting_.clear();
463 std::vector< double >& RecordCounter::_countFromDatabase_(
const IdCondSet& ids) {
466 const auto& database = _parsers_[0].data.database();
467 if (ids.empty() || database.empty() || _thread_ranges_.empty()) {
468 _last_nonDB_counting_.clear();
469 _last_nonDB_ids_.clear();
470 return _last_nonDB_counting_;
475 const auto nodeId2columns = _getNodeIds2Columns_(ids);
479 const std::size_t ids_size = ids.size();
480 std::size_t counting_vect_size = std::size_t(1);
481 std::vector< std::size_t > domain_sizes(ids_size);
482 std::vector< std::pair< std::size_t, std::size_t > > cols_offsets(ids_size);
484 std::size_t i = std::size_t(0);
485 for (
const auto id: ids) {
486 const std::size_t domain_size = database.domainSize(nodeId2columns[
id]);
487 domain_sizes[i] = domain_size;
488 cols_offsets[i].first = nodeId2columns[id];
489 cols_offsets[i].second = counting_vect_size;
490 counting_vect_size *= domain_size;
498 cols_offsets.begin(),
500 [](
const std::pair< std::size_t, std::size_t >& a,
501 const std::pair< std::size_t, std::size_t >& b) ->
bool { return a.first < b.first; });
504 const std::size_t nb_ranges = _thread_ranges_.size();
505 const auto max_nb_threads = ThreadNumberManager::getNumberOfThreads();
506 const std::size_t nb_threads = nb_ranges <= max_nb_threads ? nb_ranges : max_nb_threads;
507 while (_parsers_.size() < nb_threads) {
508 ThreadData< DBRowGeneratorParser > new_parser(_parsers_[0]);
509 _parsers_.push_back(std::move(new_parser));
515 std::vector< std::size_t > cols_of_interest(ids_size);
516 for (std::size_t i = std::size_t(0); i < ids_size; ++i) {
517 cols_of_interest[i] = cols_offsets[i].first;
519 for (
auto& parser: _parsers_) {
520 parser.data.setColumnsOfInterest(cols_of_interest);
526 std::vector< double > counting_vect(counting_vect_size, 0.0);
527 std::vector< ThreadData< std::vector< double > > > thread_countings(
529 ThreadData< std::vector< double > >(counting_vect));
533 auto threadedCount = [
this, nb_ranges, ids_size, &thread_countings, cols_offsets](
534 const std::size_t this_thread,
535 const std::size_t nb_threads,
536 const std::size_t nb_loop) ->
void {
537 if (this_thread + nb_loop < nb_ranges) {
539 DBRowGeneratorParser& parser = this->_parsers_[this_thread].data;
540 parser.setRange(this->_thread_ranges_[this_thread + nb_loop].first,
541 this->_thread_ranges_[this_thread + nb_loop].second);
542 std::vector< double >& counts = thread_countings[this_thread].data;
546 while (parser.hasRows()) {
548 const DBRow< DBTranslatedValue >& row = parser.row();
551 std::size_t offset = std::size_t(0);
552 for (std::size_t i = std::size_t(0); i < ids_size; ++i) {
553 offset += row[cols_offsets[i].first].discr_val * cols_offsets[i].second;
556 counts[offset] += row.weight();
566 for (std::size_t i = std::size_t(0); i < nb_ranges; i += nb_threads) {
567 ThreadExecutor::execute(nb_threads, threadedCount, i);
572 for (std::size_t k = std::size_t(0); k < nb_threads; ++k) {
573 const auto& thread_counting = thread_countings[k].data;
574 for (std::size_t r = std::size_t(0); r < counting_vect_size; ++r) {
575 counting_vect[r] += thread_counting[r];
581 _last_DB_counting_ = std::move(counting_vect);
583 return _last_DB_counting_;
587 void RecordCounter::_checkRanges_(
588 const std::vector< std::pair< std::size_t, std::size_t > >& new_ranges)
const {
589 const std::size_t dbsize = _parsers_[0].data.database().nbRows();
590 std::vector< std::pair< std::size_t, std::size_t > > incorrect_ranges;
591 for (
const auto& range: new_ranges) {
592 if ((range.first >= range.second) || (range.second > dbsize)) {
593 incorrect_ranges.push_back(range);
596 if (!incorrect_ranges.empty()) {
597 std::string rangeList;
599 for (
const auto& range: incorrect_ranges) {
600 if (!first) rangeList +=
", ";
602 rangeList += std::format(
"[{};{})", range.first, range.second);
604 if (incorrect_ranges.size() > 1) {
606 std::format(
"It is impossible to set the ranges because the following ones "
611 std::format(
"It is impossible to set the ranges because the following one "
619 void RecordCounter::_dispatchRangesToThreads_() {
620 _thread_ranges_.clear();
623 bool add_range =
false;
624 if (_ranges_.empty()) {
625 const auto& database = _parsers_[0].data.database();
627 std::pair< std::size_t, std::size_t >(std::size_t(0), database.nbRows()));
632 const auto max_nb_threads = ThreadNumberManager::getNumberOfThreads();
633 for (
const auto& range: _ranges_) {
634 if (range.second > range.first) {
635 const std::size_t range_size = range.second - range.first;
636 std::size_t nb_threads = range_size / _min_nb_rows_per_thread_;
637 if (nb_threads < 1) nb_threads = 1;
638 else if (nb_threads > max_nb_threads) nb_threads = max_nb_threads;
639 std::size_t nb_rows_par_thread = range_size / nb_threads;
640 std::size_t rest_rows = range_size - nb_rows_par_thread * nb_threads;
642 std::size_t begin_index = range.first;
643 for (std::size_t i = std::size_t(0); i < nb_threads; ++i) {
644 std::size_t end_index = begin_index + nb_rows_par_thread;
645 if (rest_rows != std::size_t(0)) {
649 _thread_ranges_.push_back(
650 std::pair< std::size_t, std::size_t >(begin_index, end_index));
651 begin_index = end_index;
655 if (add_range) _ranges_.clear();
661 std::sort(_thread_ranges_.begin(),
662 _thread_ranges_.end(),
663 [](
const std::pair< std::size_t, std::size_t >& a,
664 const std::pair< std::size_t, std::size_t >& b) ->
bool {
665 return (a.second - a.first) > (b.second - b.first);
670 void RecordCounter::setRanges(
671 const std::vector< std::pair< std::size_t, std::size_t > >& new_ranges) {
673 _checkRanges_(new_ranges);
676 const std::size_t new_size = new_ranges.size();
677 std::vector< std::pair< std::size_t, std::size_t > > ranges(new_size);
678 for (std::size_t i = std::size_t(0); i < new_size; ++i) {
679 ranges[i].first = new_ranges[i].first;
680 ranges[i].second = new_ranges[i].second;
684 _ranges_ = std::move(ranges);
687 _dispatchRangesToThreads_();
691 void RecordCounter::clearRanges() {
692 if (_ranges_.empty())
return;
695 _dispatchRangesToThreads_();
Exception : the element we looked for cannot be found.
Exception : out of bound.
Exception : wrong type for this operation.
the class used to read a row in the database and to transform it into a set of DBRow instances that c...
RecordCounter(const DBRowGeneratorParser &parser, const std::vector< std::pair< std::size_t, std::size_t > > &ranges, const Bijection< NodeId, std::size_t > &nodeId2columns=Bijection< NodeId, std::size_t >())
default constructor
#define GUM_ERROR(type, msg)
include the inlined functions if necessary
gum is the global namespace for all aGrUM entities
The class that computes counting of observations from the database.