25#ifndef STOCHTREE_CATEGORY_TRACKER_H_
26#define STOCHTREE_CATEGORY_TRACKER_H_
29#include <stochtree/log.h>
30#include <stochtree/meta.h>
50 observation_indices_.resize(num_observations_);
51 for (
int i = 0;
i < num_observations_;
i++) {
57 num_observations_ =
other.NumObservations();
58 observation_indices_.resize(num_observations_);
59 for (
int i = 0;
i < num_observations_;
i++) {
60 observation_indices_[
i] =
other.GetCategoryId(
i);
74 inline int NumObservations() {
return num_observations_; }
77 std::vector<int> observation_indices_;
84class CategorySampleTracker {
86 CategorySampleTracker(
const std::vector<int>&
group_indices) {
88 indices_ = std::vector<data_size_t>(n);
89 std::iota(indices_.begin(), indices_.end(), 0);
92 std::stable_sort(indices_.begin(), indices_.end(),
comp_op);
96 for (
int i = 0;
i < n;
i++) {
102 category_id_map_.insert({
group_indices[indices_[
i]], category_count_});
104 node_index_vector_.emplace_back();
106 category_begin_.push_back(
i);
108 category_begin_.push_back(
i);
119 node_index_vector_[category_count_ - 1].emplace_back(indices_[
i]);
125 indices_ = std::vector<data_size_t>(n);
126 std::iota(indices_.begin(), indices_.end(), 0);
129 std::stable_sort(indices_.begin(), indices_.end(),
comp_op);
133 for (
int i = 0;
i < n;
i++) {
139 category_id_map_.insert({
group_indices[indices_[
i]], category_count_});
141 node_index_vector_.emplace_back();
143 category_begin_.push_back(
i);
145 category_begin_.push_back(
i);
156 node_index_vector_[category_count_ - 1].emplace_back(indices_[
i]);
171 return category_begin_[
id] + category_length_[
id];
176 return category_length_[category_id_map_[
category_id]];
180 inline data_size_t NumCategories() {
return category_count_; }
183 std::vector<data_size_t> indices_;
186 std::vector<data_size_t>& NodeIndices(
int category_id) {
188 return node_index_vector_[
id];
197 std::map<int, int>& GetLabelMap() {
return category_id_map_; }
199 std::vector<int>& GetUniqueGroupIds() {
return unique_category_ids_; }
203 std::vector<data_size_t> category_begin_;
204 std::vector<data_size_t> category_length_;
205 std::map<int, int> category_id_map_;
206 std::vector<int> unique_category_ids_;
207 std::vector<std::vector<data_size_t>> node_index_vector_;
Class storing sample-node map for each tree in an ensemble TODO: Add run-time checks for categories w...
Definition category_tracker.h:41
static void GFRSampleOneIter(TreeEnsemble &active_forest, ForestTracker &tracker, ForestContainer &forests, LeafModel &leaf_model, ForestDataset &dataset, ColumnVector &residual, TreePrior &tree_prior, std::mt19937 &gen, std::vector< double > &variable_weights, std::vector< int > &sweep_update_indices, double global_variance, std::vector< FeatureType > &feature_types, int cutpoint_grid_size, bool keep_forest, bool pre_initialized, bool backfitting, int num_features_subsample, int num_threads, LeafSuffStatConstructorArgs &... leaf_suff_stat_args)
Definition tree_sampler.h:816
A collection of random number generation utilities.
Definition bart.h:15