StochTree 0.5.0.9000
Loading...
Searching...
No Matches
linear_regression.h
1
5#ifndef STOCHTREE_REGRESSION_H_
6#define STOCHTREE_REGRESSION_H_
7
8#include <Eigen/Dense>
9#include <stochtree/distributions.h>
10#include <stochtree/normal_sampler.h>
11
12#include <random>
13
14namespace StochTree {
15
27static double sample_univariate_gaussian_regression_coefficient(double* y, double* x, double error_variance, double prior_variance, int n, std::mt19937& gen) {
28 double sum_xx = 0.0;
29 double sum_yx = 0.0;
30 for (int i = 0; i < n; i++) {
31 sum_xx += x[i] * x[i];
32 sum_yx += y[i] * x[i];
33 }
36 return sample_standard_normal(post_mean, std::sqrt(post_var), gen);
37}
38
53static void sample_general_bivariate_gaussian_regression_coefficients(double* output, double* y, double* x1, double* x2, double error_variance, double prior_variance_11, double prior_variance_12, double prior_variance_22, int n, std::mt19937& gen) {
58 double sum_x1x1 = 0.0;
59 double sum_x1x2 = 0.0;
60 double sum_x2x2 = 0.0;
61 double sum_yx1 = 0.0;
62 double sum_yx2 = 0.0;
63 for (int i = 0; i < n; i++) {
64 sum_x1x1 += x1[i] * x1[i];
65 sum_x1x2 += x1[i] * x2[i];
66 sum_x2x2 += x2[i] * x2[i];
67 sum_yx1 += y[i] * x1[i];
68 sum_yx2 += y[i] * x2[i];
69 }
79 double chol_var_11 = std::sqrt(post_var_11);
81 double chol_var_22 = std::sqrt(post_var_22 - chol_var_12 * chol_var_12);
82 double z1 = sample_standard_normal(0.0, 1.0, gen);
83 double z2 = sample_standard_normal(0.0, 1.0, gen);
86}
87
101static void sample_diagonal_bivariate_gaussian_regression_coefficients(double* output, double* y, double* x1, double* x2, double error_variance, double prior_variance_11, double prior_variance_22, int n, std::mt19937& gen) {
102 double inv_prior_var_11 = 1.0 / prior_variance_11;
103 double inv_prior_var_22 = 1.0 / prior_variance_22;
104 double sum_x1x1 = 0.0;
105 double sum_x1x2 = 0.0;
106 double sum_x2x2 = 0.0;
107 double sum_yx1 = 0.0;
108 double sum_yx2 = 0.0;
109 for (int i = 0; i < n; i++) {
110 sum_x1x1 += x1[i] * x1[i];
111 sum_x1x2 += x1[i] * x2[i];
112 sum_x2x2 += x2[i] * x2[i];
113 sum_yx1 += y[i] * x1[i];
114 sum_yx2 += y[i] * x2[i];
115 }
125 double chol_var_11 = std::sqrt(post_var_11);
127 double chol_var_22 = std::sqrt(post_var_22 - chol_var_12 * chol_var_12);
128 double z1 = sample_standard_normal(0.0, 1.0, gen);
129 double z2 = sample_standard_normal(0.0, 1.0, gen);
132}
133
144static Eigen::VectorXd sample_general_gaussian_regression_coefficients(const Eigen::Ref<const Eigen::VectorXd>& y, const Eigen::Ref<const Eigen::MatrixXd>& X, double error_variance, const Eigen::Ref<const Eigen::MatrixXd>& prior_variance, int n, std::mt19937& gen) {
145 int p = X.cols();
146 Eigen::MatrixXd inv_prior_var = prior_variance.inverse();
147 Eigen::MatrixXd XtX = X.transpose() * X;
148 Eigen::VectorXd Xty = X.transpose() * y;
149 Eigen::MatrixXd post_var_pre_inv = inv_prior_var + XtX / error_variance;
150 Eigen::MatrixXd post_var = post_var_pre_inv.inverse();
151 Eigen::VectorXd post_mean = post_var * (Xty / error_variance);
152 Eigen::LLT<Eigen::MatrixXd> chol(post_var);
153 Eigen::MatrixXd L = chol.matrixL();
154 Eigen::VectorXd z(p);
155 for (int i = 0; i < p; i++) {
156 z(i) = sample_standard_normal(0.0, 1.0, gen);
157 }
158 return post_mean + L * z;
159}
160
161} // namespace StochTree
162
163#endif // STOCHTREE_REGRESSION_H_
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
static double sample_univariate_gaussian_regression_coefficient(double *y, double *x, double error_variance, double prior_variance, int n, std::mt19937 &gen)
Sample a regression coefficient from the posterior distribution of a univariate Gaussian regression m...
Definition linear_regression.h:27
double sample_standard_normal(double mean, double sd, std::mt19937 &gen)
Definition distributions.h:94
static void sample_general_bivariate_gaussian_regression_coefficients(double *output, double *y, double *x1, double *x2, double error_variance, double prior_variance_11, double prior_variance_12, double prior_variance_22, int n, std::mt19937 &gen)
Sample regression coefficients from the posterior distribution of a bivariate Gaussian regression mod...
Definition linear_regression.h:53
static Eigen::VectorXd sample_general_gaussian_regression_coefficients(const Eigen::Ref< const Eigen::VectorXd > &y, const Eigen::Ref< const Eigen::MatrixXd > &X, double error_variance, const Eigen::Ref< const Eigen::MatrixXd > &prior_variance, int n, std::mt19937 &gen)
Sample regression coefficients from the posterior distribution of a bivariate Gaussian regression mod...
Definition linear_regression.h:144
static void sample_diagonal_bivariate_gaussian_regression_coefficients(double *output, double *y, double *x1, double *x2, double error_variance, double prior_variance_11, double prior_variance_22, int n, std::mt19937 &gen)
Sample regression coefficients from the posterior distribution of a bivariate Gaussian regression mod...
Definition linear_regression.h:101