5#ifndef STOCHTREE_PREDICTION_H_
6#define STOCHTREE_PREDICTION_H_
9#include <stochtree/bart.h>
10#include <stochtree/bcf.h>
11#include <stochtree/container.h>
12#include <stochtree/meta.h>
13#include <stochtree/random_effects.h>
42 bool mean_forest =
false;
43 bool variance_forest =
false;
44 bool random_effects =
false;
57 std::vector<double> y_hat;
60 std::vector<double> mean_forest_predictions;
63 std::vector<double> variance_forest_predictions;
66 std::vector<double> rfx_predictions;
80 bool has_variance_forest =
false;
82 BARTRFXModelSpec rfx_model_spec;
83 PredType pred_type = PredType::kPosterior;
85 PredScale pred_scale = PredScale::kLinear;
88 int cloglog_num_classes = 0;
109 bool prognostic_function =
false;
111 bool conditional_variance =
false;
112 bool random_effects =
false;
125 std::vector<double> y_hat;
128 std::vector<double> mu_x;
131 std::vector<double> tau_x;
135 std::vector<double> prognostic_function;
139 std::vector<double> cate;
142 std::vector<double> conditional_variance;
145 std::vector<double> random_effects;
156 int treatment_dim = 0;
159 bool has_variance_forest =
false;
160 bool has_rfx =
false;
161 BCFRFXModelSpec rfx_model_spec;
162 bool adaptive_coding =
false;
163 bool sample_tau_0 =
false;
164 PredType pred_type = PredType::kPosterior;
166 PredScale pred_scale = PredScale::kLinear;
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
BARTPredictionResult predict_bart_model(BARTData &data, BARTSamples &samples, BARTPredictionMetadata &metadata)
BART prediction function.
PredType
Determines whether posterior predictions are returned as-is or pre-aggregated.
Definition prediction.h:18
PredScale
Determines the scale of predictions (i.e. whether a probability / class transformation is applied)
Definition prediction.h:33
BCFPredictionResult predict_bcf_model(BCFData &data, BCFSamples &samples, BCFPredictionMetadata &metadata)
BCF prediction function.
Selector for model terms that should be predicted.
Definition prediction.h:40
Struct returning BART model predictions.
Definition prediction.h:55
Selector for model terms that should be predicted.
Definition prediction.h:105
Struct returning BCF model predictions.
Definition prediction.h:123