StochTree 0.5.0.9000
Loading...
Searching...
No Matches
random_effects.h
1
5#ifndef STOCHTREE_RANDOM_EFFECTS_H_
6#define STOCHTREE_RANDOM_EFFECTS_H_
7
8#include <stochtree/category_tracker.h>
9#include <stochtree/cutpoint_candidates.h>
10#include <stochtree/data.h>
11#include <stochtree/ensemble.h>
12#include <stochtree/ig_sampler.h>
13#include <stochtree/log.h>
14#include <stochtree/normal_sampler.h>
15#include <stochtree/partition_tracker.h>
16#include <stochtree/prior.h>
17#include <nlohmann/json.hpp>
18#include <Eigen/Dense>
19
20#include <fstream>
21#include <map>
22#include <memory>
23#include <random>
24#include <string>
25#include <vector>
26
27namespace StochTree {
28
30class LabelMapper;
31class MultivariateRegressionRandomEffectsModel;
32class RandomEffectsContainer;
33
36 public:
37 RandomEffectsTracker(std::vector<int>& group_indices);
40 inline data_size_t GetCategoryId(int observation_num) { return sample_category_mapper_->GetCategoryId(observation_num); }
41 inline data_size_t CategoryBegin(int category_id) { return category_sample_tracker_->CategoryBegin(category_id); }
42 inline data_size_t CategoryEnd(int category_id) { return category_sample_tracker_->CategoryEnd(category_id); }
43 inline data_size_t CategorySize(int category_id) { return category_sample_tracker_->CategorySize(category_id); }
44 inline int NumCategories() { return num_categories_; }
45 inline int CategoryNumber(int category_id) { return category_sample_tracker_->CategoryNumber(category_id); }
46 SampleCategoryMapper* GetSampleCategoryMapper() { return sample_category_mapper_.get(); }
47 CategorySampleTracker* GetCategorySampleTracker() { return category_sample_tracker_.get(); }
48 std::vector<data_size_t>::iterator UnsortedNodeBeginIterator(int category_id);
49 std::vector<data_size_t>::iterator UnsortedNodeEndIterator(int category_id);
50 std::map<int, int>& GetLabelMap() { return category_sample_tracker_->GetLabelMap(); }
51 std::vector<int>& GetUniqueGroupIds() { return category_sample_tracker_->GetUniqueGroupIds(); }
52 std::vector<data_size_t>& NodeIndices(int category_id) { return category_sample_tracker_->NodeIndices(category_id); }
53 std::vector<data_size_t>& NodeIndicesInternalIndex(int internal_category_id) { return category_sample_tracker_->NodeIndicesInternalIndex(internal_category_id); }
54 double* GetPredictions() { return rfx_predictions_.data(); }
55 double GetPrediction(data_size_t observation_num) { return rfx_predictions_.at(observation_num); }
56 void SetPrediction(data_size_t observation_num, double pred) { rfx_predictions_.at(observation_num) = pred; }
65
66 private:
68 std::unique_ptr<SampleCategoryMapper> sample_category_mapper_;
70 std::unique_ptr<CategorySampleTracker> category_sample_tracker_;
72 std::vector<double> rfx_predictions_;
74 int num_categories_;
75 int num_observations_;
76};
77
80 public:
81 LabelMapper() {}
82 LabelMapper(std::map<int, int> label_map) {
83 label_map_ = label_map;
84 for (const auto& [key, value] : label_map) keys_.push_back(key);
85 }
86 ~LabelMapper() {}
87 void LoadFromLabelMap(std::map<int, int> label_map) {
88 label_map_ = label_map;
89 for (const auto& [key, value] : label_map) keys_.push_back(key);
90 }
91 bool ContainsLabel(int category_id) {
92 auto pos = label_map_.find(category_id);
93 return pos != label_map_.end();
94 }
95 int CategoryNumber(int category_id) {
96 return label_map_[category_id];
97 }
98 void CopyFromOther(LabelMapper& other) {
99 keys_ = other.Keys();
100 label_map_ = other.Map();
101 }
102 void SaveToJsonFile(std::string filename) {
103 nlohmann::json model_json = this->to_json();
104 std::ofstream output_file(filename);
105 output_file << model_json << std::endl;
106 }
107 void LoadFromJsonFile(std::string filename) {
108 std::ifstream f(filename);
109 nlohmann::json rfx_label_mapper_json = nlohmann::json::parse(f);
110 this->Reset();
111 this->from_json(rfx_label_mapper_json);
112 }
113 std::string DumpJsonString() {
114 nlohmann::json model_json = this->to_json();
115 return model_json.dump();
116 }
117 void LoadFromJsonString(std::string& json_string) {
118 nlohmann::json rfx_label_mapper_json = nlohmann::json::parse(json_string);
119 this->Reset();
120 this->from_json(rfx_label_mapper_json);
121 }
122 std::vector<int>& Keys() { return keys_; }
123 std::map<int, int>& Map() { return label_map_; }
124 void Reset() {
125 label_map_.clear();
126 keys_.clear();
127 }
128 nlohmann::json to_json();
129 void from_json(const nlohmann::json& rfx_label_mapper_json);
130
131 private:
132 std::map<int, int> label_map_;
133 std::vector<int> keys_;
134};
135
138 public:
140 normal_sampler_ = MultivariateNormalSampler();
141 ig_sampler_ = InverseGammaSampler();
142 num_components_ = num_components;
143 num_groups_ = num_groups;
144 working_parameter_ = Eigen::VectorXd(num_components_);
145 group_parameters_ = Eigen::MatrixXd(num_components_, num_groups_);
146 group_parameter_covariance_ = Eigen::MatrixXd(num_components_, num_components_);
147 working_parameter_covariance_ = Eigen::MatrixXd(num_components_, num_components_);
148 }
150
153
156 void SampleWorkingParameter(RandomEffectsDataset& dataset, ColumnVector& residual, RandomEffectsTracker& tracker, double global_variance, std::mt19937& gen);
157 void SampleGroupParameters(RandomEffectsDataset& dataset, ColumnVector& residual, RandomEffectsTracker& tracker, double global_variance, std::mt19937& gen);
158 void SampleVarianceComponents(RandomEffectsDataset& dataset, ColumnVector& residual, RandomEffectsTracker& tracker, double global_variance, std::mt19937& gen);
159
161 void SetWorkingParameter(Eigen::VectorXd& working_parameter) {
162 working_parameter_ = working_parameter;
163 }
164 void SetGroupParameters(Eigen::MatrixXd& group_parameters) {
165 group_parameters_ = group_parameters;
166 }
167 void SetGroupParameter(Eigen::VectorXd& group_parameter, int group_id) {
168 group_parameters_(Eigen::all, group_id) = group_parameter;
169 }
170 void SetWorkingParameterCovariance(Eigen::MatrixXd& working_parameter_covariance) {
171 working_parameter_covariance_ = working_parameter_covariance;
172 }
173 void SetGroupParameterCovariance(Eigen::MatrixXd& group_parameter_covariance) {
174 group_parameter_covariance_ = group_parameter_covariance;
175 }
176 void SetGroupParameterVarianceComponent(double value, int component_id) {
177 group_parameter_covariance_(component_id, component_id) = value;
178 }
179 void SetVariancePriorShape(double value) {
180 variance_prior_shape_ = value;
181 }
182 void SetVariancePriorScale(double value) {
183 variance_prior_scale_ = value;
184 }
185
187 Eigen::VectorXd& GetWorkingParameter() {
188 return working_parameter_;
189 }
190 Eigen::MatrixXd& GetGroupParameters() {
191 return group_parameters_;
192 }
193 Eigen::MatrixXd& GetWorkingParameterCovariance() {
194 return working_parameter_covariance_;
195 }
196 Eigen::MatrixXd& GetGroupParameterCovariance() {
197 return group_parameter_covariance_;
198 }
199 double GetVariancePriorShape() {
200 return variance_prior_shape_;
201 }
202 double GetVariancePriorScale() {
203 return variance_prior_scale_;
204 }
205 inline int NumComponents() { return num_components_; }
206 inline int NumGroups() { return num_groups_; }
207
208 std::vector<double> Predict(RandomEffectsDataset& dataset, RandomEffectsTracker& tracker) {
209 std::vector<double> output(dataset.NumObservations());
210 PredictInplace(dataset, tracker, output);
211 return output;
212 }
213
214 void PredictInplace(RandomEffectsDataset& dataset, RandomEffectsTracker& tracker, std::vector<double>& output) {
215 Eigen::MatrixXd X = dataset.GetBasis();
216 std::vector<int> group_labels = dataset.GetGroupLabels();
217 CHECK_EQ(X.rows(), group_labels.size());
218 int n = X.rows();
219 CHECK_EQ(n, output.size());
220 Eigen::MatrixXd alpha_diag = working_parameter_.asDiagonal().toDenseMatrix();
221 int group_ind;
222 for (int i = 0; i < n; i++) {
223 group_ind = tracker.CategoryNumber(group_labels[i]);
224 output[i] = X(i, Eigen::all) * alpha_diag * group_parameters_(Eigen::all, group_ind);
225 }
226 }
227
228 void AddCurrentPredictionToResidual(RandomEffectsDataset& dataset, RandomEffectsTracker& tracker, ColumnVector& residual) {
229 data_size_t n = dataset.NumObservations();
230 CHECK_EQ(n, residual.NumRows());
231 double current_pred;
232 double new_resid;
233 for (data_size_t i = 0; i < n; i++) {
234 current_pred = tracker.GetPrediction(i);
235 new_resid = residual.GetElement(i) + current_pred;
236 residual.SetElement(i, new_resid);
237 }
238 }
239
240 void SubtractNewPredictionFromResidual(RandomEffectsDataset& dataset, RandomEffectsTracker& tracker, ColumnVector& residual) {
241 Eigen::MatrixXd X = dataset.GetBasis();
242 std::vector<int> group_labels = dataset.GetGroupLabels();
243 CHECK_EQ(X.rows(), group_labels.size());
244 int n = X.rows();
245 double new_pred;
246 double new_resid;
247 Eigen::MatrixXd alpha_diag = working_parameter_.asDiagonal().toDenseMatrix();
248 int group_ind;
249 for (int i = 0; i < n; i++) {
250 group_ind = tracker.CategoryNumber(group_labels[i]);
251 new_pred = X(i, Eigen::all) * alpha_diag * group_parameters_(Eigen::all, group_ind);
252 new_resid = residual.GetElement(i) - new_pred;
253 residual.SetElement(i, new_resid);
254 tracker.SetPrediction(i, new_pred);
255 }
256 }
257
270
271 private:
273 MultivariateNormalSampler normal_sampler_;
274 InverseGammaSampler ig_sampler_;
275
277 int num_components_;
278 int num_groups_;
279
283 Eigen::VectorXd working_parameter_;
284 Eigen::MatrixXd group_parameters_;
285
287 Eigen::MatrixXd group_parameter_covariance_;
288
290 Eigen::MatrixXd working_parameter_covariance_;
291
293 double variance_prior_shape_;
294 double variance_prior_scale_;
295};
296
298 public:
300 num_components_ = num_components;
301 num_groups_ = num_groups;
302 num_samples_ = 0;
303 }
305 num_components_ = 0;
306 num_groups_ = 0;
307 num_samples_ = 0;
308 }
310 void SaveToJsonFile(std::string filename) {
311 nlohmann::json model_json = this->to_json();
312 std::ofstream output_file(filename);
313 output_file << model_json << std::endl;
314 }
315 void LoadFromJsonFile(std::string filename) {
316 std::ifstream f(filename);
317 nlohmann::json rfx_container_json = nlohmann::json::parse(f);
318 this->Reset();
319 this->from_json(rfx_container_json);
320 }
321 std::string DumpJsonString() {
322 nlohmann::json model_json = this->to_json();
323 return model_json.dump();
324 }
325 void LoadFromJsonString(std::string& json_string) {
326 nlohmann::json rfx_container_json = nlohmann::json::parse(json_string);
327 this->Reset();
328 this->from_json(rfx_container_json);
329 }
330 void CopyFromOther(RandomEffectsContainer& other) {
331 this->Reset();
332 num_samples_ = other.NumSamples();
333 num_components_ = other.NumComponents();
334 num_groups_ = other.NumGroups();
335 beta_ = other.GetBeta();
336 alpha_ = other.GetAlpha();
337 xi_ = other.GetXi();
338 sigma_xi_ = other.GetSigma();
339 }
341 void DeleteSample(int sample_num);
342 void Predict(RandomEffectsDataset& dataset, LabelMapper& label_mapper, std::vector<double>& output);
343 inline int NumSamples() { return num_samples_; }
344 inline int NumComponents() { return num_components_; }
345 inline int NumGroups() { return num_groups_; }
346 inline void SetNumSamples(int num_samples) { num_samples_ = num_samples; }
347 inline void SetNumComponents(int num_components) { num_components_ = num_components; }
348 inline void SetNumGroups(int num_groups) { num_groups_ = num_groups; }
349 void Reset() {
350 num_samples_ = 0;
351 num_components_ = 0;
352 num_groups_ = 0;
353 beta_.clear();
354 alpha_.clear();
355 xi_.clear();
356 sigma_xi_.clear();
357 }
358 std::vector<double>& GetBeta() { return beta_; }
359 std::vector<double>& GetAlpha() { return alpha_; }
360 std::vector<double>& GetXi() { return xi_; }
361 std::vector<double>& GetSigma() { return sigma_xi_; }
362 nlohmann::json to_json();
363 void from_json(const nlohmann::json& rfx_container_json);
364 void append_from_json(const nlohmann::json& rfx_container_json);
365
366 private:
367 int num_samples_;
368 int num_components_;
369 int num_groups_;
370 std::vector<double> beta_;
371 std::vector<double> alpha_;
372 std::vector<double> xi_;
373 std::vector<double> sigma_xi_;
377};
378
384inline void AppendRandomEffectsContainerSamples(std::unique_ptr<RandomEffectsContainer>& dst,
385 const std::unique_ptr<RandomEffectsContainer>& src) {
386 if (src == nullptr && dst == nullptr) return;
387 if (src == nullptr || dst == nullptr) {
388 Log::Fatal("Cannot merge samples: random effects container present in one chain but not the other");
389 }
390 // Check that group and basis dimensions match between the two containers before appending
391 if (src->NumComponents() != dst->NumComponents()) {
392 Log::Fatal("Cannot merge samples: random effects container has %d components in one chain but %d in the other", src->NumComponents(), dst->NumComponents());
393 }
394 if (src->NumGroups() != dst->NumGroups()) {
395 Log::Fatal("Cannot merge samples: random effects container has %d groups in one chain but %d in the other", src->NumGroups(), dst->NumGroups());
396 }
397 dst->SetNumSamples(dst->NumSamples() + src->NumSamples());
398 std::vector<double>& dst_beta = dst->GetBeta();
399 std::vector<double>& dst_alpha = dst->GetAlpha();
400 std::vector<double>& dst_xi = dst->GetXi();
401 std::vector<double>& dst_sigma_xi = dst->GetSigma();
402 std::vector<double>& src_beta = src->GetBeta();
403 std::vector<double>& src_alpha = src->GetAlpha();
404 std::vector<double>& src_xi = src->GetXi();
405 std::vector<double>& src_sigma_xi = src->GetSigma();
406 dst_beta.insert(dst_beta.end(), src_beta.begin(), src_beta.end());
407 dst_alpha.insert(dst_alpha.end(), src_alpha.begin(), src_alpha.end());
408 dst_xi.insert(dst_xi.end(), src_xi.begin(), src_xi.end());
409 dst_sigma_xi.insert(dst_sigma_xi.end(), src_sigma_xi.begin(), src_sigma_xi.end());
410}
411
412} // namespace StochTree
413
414#endif // STOCHTREE_RANDOM_EFFECTS_H_
Internal wrapper around Eigen::VectorXd interface for univariate floating point data....
Definition data.h:193
Definition ig_sampler.h:10
Standalone container for the map from category IDs to 0-based indices.
Definition random_effects.h:79
Definition normal_sampler.h:26
Posterior computation and sampling and state storage for random effects model with a group-level mult...
Definition random_effects.h:137
double VarianceComponentShape(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &rfx_tracker, double global_variance, int component_id)
Compute the posterior shape of the group variance component, conditional on the working and group par...
void ResetFromSample(RandomEffectsContainer &rfx_container, int sample_num)
Reconstruction from serialized model parameter samples.
Eigen::VectorXd & GetWorkingParameter()
Getters.
Definition random_effects.h:187
void SampleRandomEffects(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &tracker, double global_variance, std::mt19937 &gen)
Samplers.
Eigen::VectorXd WorkingParameterMean(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &rfx_tracker, double global_variance)
Compute the posterior mean of the working parameter, conditional on the group parameters and the vari...
Eigen::VectorXd GroupParameterMean(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &rfx_tracker, double global_variance, int group_id)
Compute the posterior mean of a group parameter, conditional on the working parameter and the varianc...
Eigen::MatrixXd GroupParameterVariance(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &rfx_tracker, double global_variance, int group_id)
Compute the posterior covariance of a group parameter, conditional on the working parameter and the v...
Eigen::MatrixXd WorkingParameterVariance(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &rfx_tracker, double global_variance)
Compute the posterior covariance of the working parameter, conditional on the group parameters and th...
void SetWorkingParameter(Eigen::VectorXd &working_parameter)
Setters.
Definition random_effects.h:161
double VarianceComponentScale(RandomEffectsDataset &dataset, ColumnVector &residual, RandomEffectsTracker &rfx_tracker, double global_variance, int component_id)
Compute the posterior scale of the group variance component, conditional on the working and group par...
Definition random_effects.h:297
API for loading and accessing data used to sample (additive) random effects.
Definition data.h:523
Wrapper around data structures for random effects sampling algorithms.
Definition random_effects.h:35
void RootReset(MultivariateRegressionRandomEffectsModel &rfx_model, RandomEffectsDataset &rfx_dataset, ColumnVector &residual)
Resets RFX tracker to initial default. Assumes tracker already exists in main memory....
void ResetFromSample(MultivariateRegressionRandomEffectsModel &rfx_model, RandomEffectsDataset &rfx_dataset, ColumnVector &residual)
Resets RFX tracker based on a specific sample. Assumes tracker already exists in main memory.
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
void AppendRandomEffectsContainerSamples(std::unique_ptr< RandomEffectsContainer > &dst, const std::unique_ptr< RandomEffectsContainer > &src)
Append every retained random effects sample from src onto the end of dst (deep copy)....
Definition random_effects.h:384