ShrinkageTrees is an R package that brings Bayesian Additive Regression Trees (BART; Chipman, George & McCulloch, 2010) to survival analysis and causal inference, with a particular focus on high-dimensional data.
The package implements BART-based models for right-censored and interval-censored survival outcomes using an accelerated failure time (AFT) formulation. Censored event times are handled through Bayesian data augmentation in the Gibbs sampler, enabling full posterior inference without proportional-hazards assumptions. For causal inference, the package provides Bayesian Causal Forests (BCF; Hahn, Murray & Carvalho, 2020), which decompose the outcome into a prognostic function \(\mu(\mathbf{x})\) and a treatment-effect function \(\tau(\mathbf{x})\), each estimated by a separate tree ensemble. This two-forest structure supports estimation of heterogeneous treatment effects (CATEs) and the average treatment effect (ATE).
A key feature is the availability of multiple regularisation strategies that can be freely combined within a single model:
| Function | Task | Prior |
|---|---|---|
HorseTrees() |
Prediction (continuous / binary / survival*) | Horseshoe |
ShrinkageTrees() |
Prediction — flexible prior choice* | Horseshoe, DART, BART, … |
SurvivalBART() |
Survival prediction* | Classical BART |
SurvivalDART() |
Sparse survival prediction* | DART (Dirichlet) |
SurvivalBCF() |
Causal survival inference* | BCF (classical) |
SurvivalShrinkageBCF() |
Sparse causal survival inference* | BCF + DART |
CausalHorseForest() |
Causal inference (all outcomes*) | Horseshoe |
CausalShrinkageForest() |
Causal inference — flexible prior* | Horseshoe, DART, BART, … |
* All survival functions support both right-censored and interval-censored outcomes.
All model-fitting functions return an S3 object with consistent
print(), summary(), predict(),
and plot() methods.
Before diving into examples, we clarify a few concepts that appear throughout the package interface.
timescale parameterEvery model-fitting function accepts an outcome_type
argument:
"continuous" — standard regression (default for most
functions)."binary" — probit BART for binary outcomes (0/1)."right-censored" — accelerated failure time model for
survival data. The outcome y contains (possibly censored)
follow-up times, and the status vector indicates events (1)
vs. censored observations (0)."interval-censored" — AFT model for interval-censored
survival data. Instead of y and status,
provide left_time and right_time vectors
specifying the lower and upper bounds of the observation window for each
individual. Three cases are distinguished:
left_time == right_time
(event observed exactly).left_time < right_time with finite
right_time (event occurred somewhere in the interval).right_time = Inf
(event not yet observed).survival::Surv(type = "interval2").For survival outcomes, the timescale argument controls
how the package treats the times:
timescale = "time" (default): the supplied values are
on the original time scale (positive numbers). The
package internally applies a log-transform, i.e. models \(\log(T) = f(\mathbf{x}) + \varepsilon\).
Predictions from summary() and predict() are
back-transformed to the time scale automatically.timescale = "log": the supplied values are
already log-transformed. No further transformation is
applied. Predictions stay on the log scale.In most applications you should use timescale = "time"
and pass the raw survival times directly.
In a BART ensemble each tree contributes a step height (leaf parameter) to the overall prediction. Classical BART assigns these step heights a fixed-variance Gaussian prior, which regularises all leaves equally. In high-dimensional settings, stronger and more adaptive regularisation is desirable. ShrinkageTrees implements two shrinkage priors that are placed directly on the step heights via a scale mixture of normals:
\[ h_\ell \mid \lambda_\ell, \tau, \omega \sim \mathcal{N}(0,\; \omega\, \lambda_\ell^2\, \tau^2). \]
Here \(\tau\) is a global shrinkage parameter shared across all leaves, \(\lambda_\ell\) is a local scale specific to leaf \(\ell\), and \(\omega\) is a fixed scaling constant. The two currently implemented instantiations are:
prior_type = "horseshoe").
Both \(\lambda_\ell\) and \(\tau\) receive independent half-Cauchy
priors. The heavy tails of the half-Cauchy allow individual leaves to
escape shrinkage when the data support a strong effect, while the global
parameter \(\tau\) pulls the bulk of
the estimates toward zero. This is the default prior in
HorseTrees() and CausalHorseForest().prior_type = "half-cauchy"). Only the local scales \(\lambda_\ell\) receive a half-Cauchy prior;
there is no global shrinkage parameter. This provides per-leaf
adaptivity without the additional pooling across the ensemble.A forest-wide variant of the horseshoe
(prior_type = "horseshoe_fw") shares a single global \(\tau\) across all trees in the forest
rather than one per tree. These priors can be selected in
ShrinkageTrees() and CausalShrinkageForest()
via the prior_type argument, and can be combined with
DART’s Dirichlet splitting prior for simultaneous structural and
parametric regularisation.
The two most important hyperparameters are local_hp and
global_hp. These control the horseshoe prior on the step
heights (leaf parameters):
\[ \mu_{jl} \mid \lambda_{jl}, \tau_j \sim \mathcal{N}(0, \lambda_{jl}^2 \tau_j^2), \qquad \lambda_{jl} \sim \text{C}^+(0, \texttt{local\_hp}), \qquad \tau_j \sim \text{C}^+(0, \texttt{global\_hp}), \]
where \(\text{C}^+\) denotes the half-Cauchy distribution. Smaller values produce stronger shrinkage toward zero; larger values allow more variation.
Only the product local_hp * global_hp
identifies the prior. The leaf standard deviation is \(\sqrt{\omega}\,\lambda_{jl}\tau_j\), and
the product of two half-Cauchy variables depends on the two scales only
through their product, so any split between local and global with the
same product gives the same prior. Setting them equal therefore loses
nothing, and is what all four functions do by default.
HorseTrees() and
CausalHorseForest() provide a convenience
parameter k that sets both scales:
local_hp = global_hp = k / sqrt(number_of_trees).
ShrinkageTrees() and
CausalShrinkageForest() expose
local_hp and global_hp directly. Leaving them
NULL applies the same default.
The defaults are:
| function | default | useful range |
|---|---|---|
HorseTrees(), ShrinkageTrees() |
k = 1.0 |
roughly 0.5 to 1.5 |
CausalHorseForest(),
CausalShrinkageForest() |
k = 1.5 |
1 to 2 |
These were calibrated over a range of simulated settings and gave
near-nominal pointwise coverage for p between 50 and 5000.
Within each range:
k shrinks more aggressively —
tighter credible intervals, more leaves pulled to zero. Better when the
signal is sparse, the dimension is high, the sample is small, or (for
the causal models) the treatment effect is close to homogeneous.k is more conservative — wider
intervals and more freedom for individual leaves. Better when the
function is genuinely complex and the sample supports estimating
it.The defaults sit in the middle of each range, and in our simulations the penalty for being anywhere inside it was small.
The causal range sits higher because the treatment forest enters the model as \(b\,\tau(x)\) with \(b = \pm 1/2\), halving its contribution to the response. Measured on the scale of that contribution the two recommendations nearly coincide.
Note that k scales the product of the two
half-Cauchy scales, so the implied prior scale of the ensemble grows
with \(k^2\) — doubling k
loosens the prior fourfold. Move it in modest steps.
The defaults are a starting point, not a substitute for tuning.
k interacts with the signal-to-noise ratio, the sparsity of
the true function and the censoring rate, none of which the default can
know about, so selecting k from the data is
preferable whenever the computational budget allows.
Because k is a single scalar, a coarse grid search with
K-fold cross-validation is usually enough. For survival
outcomes the concordance index on held-out folds is a natural
criterion:
library(survival)
k_grid <- c(0.5, 0.75, 1.0, 1.25, 1.5)
folds <- sample(rep(1:5, length.out = length(y)))
cv_score <- function(k) {
scores <- vapply(1:5, function(f) {
te <- which(folds == f); tr <- setdiff(seq_along(y), te)
fit <- HorseTrees(
y = y[tr], status = status[tr], X_train = X[tr, ], X_test = X[te, ],
outcome_type = "right-censored", timescale = "log",
k = k, N_post = 1000, N_burn = 1000, verbose = FALSE
)
concordance(Surv(exp(y[te]), status[te]) ~ exp(fit$test_predictions))$concordance
}, numeric(1))
mean(scores)
}
scores <- vapply(k_grid, cv_score, numeric(1))
k_best <- k_grid[which.max(scores)]Two cautions. Check that the selected value is interior to the grid: a winner at either end means the grid did not bracket the optimum and should be widened. And prefer a criterion that can actually be optimised — measures of predictive discrimination such as the concordance index often improve monotonically as the prior loosens, in which case they will simply select the largest value on offer. Criteria that penalise over-dispersed posteriors, such as held-out log predictive density, or a direct check of interval coverage against the nominal level, are more informative for choosing a shrinkage scale.
The survival functions (SurvivalBART,
SurvivalDART) use k to calibrate the standard
BART leaf prior:
local_hp = range(log(y)) / (2 * k * sqrt(number_of_trees)).
The default k = 2 follows Chipman et al. (2010). This is a
different parameterisation and is unaffected by the change described
above.
store_posterior_sample flagWhen store_posterior_sample = TRUE, the fitted object
stores the full \(N_\text{post} \times
n\) matrix of posterior draws for predictions. This is needed
for:
predict() on new data (it re-runs the sampler
internally, so posterior samples are always produced);plot(fit, type = "ate") and
plot(fit, type = "cate"), which require the full posterior
distribution;plot(fit, type = "survival") — full posterior credible
bands over both \(\mu_i\) and \(\sigma\) (without samples, only
sigma-uncertainty bands are available);When FALSE, only posterior means and \(\sigma\) draws are stored, saving memory.
print() and summary() work in both cases, but
predict() will not be available.
treatment_coding)All causal model functions — CausalHorseForest(),
CausalShrinkageForest(), SurvivalBCF(), and
SurvivalShrinkageBCF() — decompose the outcome as \[
Y_i = \mu(\mathbf{x}_i) + b_i \cdot \tau(\mathbf{x}_i) + \varepsilon_i,
\] where \(b_i\) is a scalar
that depends on the treatment assignment \(Z_i\). The treatment_coding
argument controls how \(b_i\) is
defined. Four options are available:
"centered" (default). \(b_i = Z_i - 1/2\), so that \(b_i \in \{-1/2,\; 1/2\}\). This is the
original BCF parameterisation.
"binary". \(b_i = Z_i\), so that \(b_i \in \{0,\; 1\}\). Standard binary
coding; the treatment forest captures the full effect of treatment on
the treated.
"adaptive". \(b_i = Z_i - \hat{e}(\mathbf{x}_i)\), where
\(\hat{e}(\mathbf{x}_i)\) is the
estimated propensity score. This follows Hahn, Murray & Carvalho
(2020) and is the coding used in the bcf R package. When
using this option, a propensity vector must be
supplied.
"invariant". Parameter-expanded
(invariant) treatment coding. The coding parameters \(b_0\) and \(b_1\) are assigned \(N(0,\; 1/2)\) priors and estimated within
the Gibbs sampler via conjugate normal updates: \[
Y_i = \mu(\mathbf{x}_i) + b_{Z_i} \cdot \tilde{\tau}(\mathbf{x}_i) +
\varepsilon_i,
\qquad b_0,\; b_1 \sim N(0,\; 1/2).
\] The treatment effect is \(\tau(\mathbf{x}_i) = (b_1 - b_0) \cdot
\tilde{\tau}(\mathbf{x}_i)\), and the posterior draws of \(b_0\) and \(b_1\) are returned in the fitted object.
This parameterisation is invariant to the coding of the treatment
indicator (Hahn et al., 2020, Section 5.2).
The examples below illustrate each option on a simple continuous-outcome causal model.
set.seed(50)
n_tc <- 60; p_tc <- 5
X_tc <- matrix(rnorm(n_tc * p_tc), n_tc, p_tc)
W_tc <- rbinom(n_tc, 1, 0.5)
tau_tc <- 1.5 * (X_tc[, 1] > 0)
y_tc <- X_tc[, 1] + W_tc * tau_tc + rnorm(n_tc, sd = 0.5)# Centered (default)
fit_tc_cen <- CausalHorseForest(
y = y_tc,
X_train_control = X_tc, X_train_treat = X_tc,
treatment_indicator_train = W_tc,
treatment_coding = "centered",
number_of_trees = 5, N_post = 50, N_burn = 25,
store_posterior_sample = TRUE, verbose = FALSE
)
cat("Centered — ATE:",
round(mean(fit_tc_cen$train_predictions_treat), 3), "\n")
#> Centered — ATE: 0.887# Binary
fit_tc_bin <- CausalHorseForest(
y = y_tc,
X_train_control = X_tc, X_train_treat = X_tc,
treatment_indicator_train = W_tc,
treatment_coding = "binary",
number_of_trees = 5, N_post = 50, N_burn = 25,
store_posterior_sample = TRUE, verbose = FALSE
)
cat("Binary — ATE:",
round(mean(fit_tc_bin$train_predictions_treat), 3), "\n")
#> Binary — ATE: 0.942# Adaptive (requires propensity scores)
ps_tc <- pnorm(0.3 * X_tc[, 1]) # simple propensity model for illustration
fit_tc_ada <- CausalHorseForest(
y = y_tc,
X_train_control = X_tc, X_train_treat = X_tc,
treatment_indicator_train = W_tc,
treatment_coding = "adaptive",
propensity = ps_tc,
number_of_trees = 5, N_post = 50, N_burn = 25,
store_posterior_sample = TRUE, verbose = FALSE
)
cat("Adaptive — ATE:",
round(mean(fit_tc_ada$train_predictions_treat), 3), "\n")
#> Adaptive — ATE: 0.528# Invariant (parameter-expanded)
fit_tc_inv <- CausalHorseForest(
y = y_tc,
X_train_control = X_tc, X_train_treat = X_tc,
treatment_indicator_train = W_tc,
treatment_coding = "invariant",
number_of_trees = 5, N_post = 50, N_burn = 25,
store_posterior_sample = TRUE, verbose = FALSE
)
cat("Invariant — ATE:",
round(mean(fit_tc_inv$train_predictions_treat), 3), "\n")
#> Invariant — ATE: 0.736
# Posterior draws of b0 and b1 are stored in the fitted object
cat("b0 posterior mean:", round(mean(fit_tc_inv$b0), 3), "\n")
#> b0 posterior mean: 0.044
cat("b1 posterior mean:", round(mean(fit_tc_inv$b1), 3), "\n")
#> b1 posterior mean: 0.348The survival functions inherit treatment_coding support.
For example,
SurvivalBCF(..., treatment_coding = "invariant") works out
of the box.
The package ships with three datasets for high-dimensional survival analysis and causal inference:
pdac — TCGA pancreatic ductal
adenocarcinoma (PAAD) cohort (n = 130). A data frame with overall
survival times, a binary treatment indicator (radiation therapy
vs. control), 14 clinical covariates, and expression values of ~3,000
genes selected by median absolute deviation.ovarian — a semi-synthetic cohort (n =
357) built on covariates from the TCGA ovarian cancer (OV) study. A
single data frame of 1,004 columns: OS_time,
OS_event, treatment, four clinical covariates,
and 997 gene expression columns named by versioned Ensembl identifiers.
Treatment assignment and outcomes are simulated from a known
data-generating process. See ?ovarian for details.ovarian_truth — the quantities behind
that process, one row per patient: the prognostic surface, the true
treatment effect, the propensity score, and the uncensored event time.
This is what makes ovarian usable for validation: estimates
can be compared against the truth rather than against the observed
outcome. See ?ovarian_truth.The pdac dataset contains overall survival times, a
binary treatment indicator (radiation therapy vs. control), clinical
covariates, and expression values of approximately 3,000 genes selected
by median absolute deviation.
library(ShrinkageTrees)
data("pdac")
# Dimensions and column overview
cat("Patients:", nrow(pdac), "\n")
#> Patients: 130
cat("Columns :", ncol(pdac), "\n")
#> Columns : 3032
cat("Clinical columns:", paste(names(pdac)[1:14], collapse = ", "), "\n")
#> Clinical columns: time, status, treatment, age, sex, grade, tumor.cellularity, tumor.purity, absolute.purity, moffitt.cluster, meth.leukocyte.percent, meth.purity.mode, stage, lymph.nodes
cat("Survival: time (months), censoring rate =",
round(1 - mean(pdac$status), 2), "\n")
#> Survival: time (months), censoring rate = 0.47
cat("Treatment: radiation =", sum(pdac$treatment),
"/ control =", sum(1 - pdac$treatment), "\n")
#> Treatment: radiation = 36 / control = 94We separate the outcome, treatment, and covariate matrix for the analyses below.
This section demonstrates the single-forest models in ShrinkageTrees. These models estimate a single function \(f(\mathbf{x})\) of the covariates, applicable to continuous, binary, and survival outcomes. We begin with a binary-outcome example that will also serve as the propensity score model for the causal analyses later.
Before fitting causal models we estimate propensity scores \(\hat{e}(\mathbf{x}) = P(W=1 \mid
\mathbf{x})\) using HorseTrees() with a binary
outcome. The probit link is used internally: predictions are on the
latent Gaussian scale and can be converted to probabilities with
pnorm().
The code block below uses reduced MCMC settings for illustration. A
real analysis would use N_post = 5000, N_burn = 5000.
ps_fit <- HorseTrees(
y = treatment,
X_train = X,
outcome_type = "binary",
k = 1.0,
N_post = 5000,
N_burn = 5000,
verbose = FALSE
)
propensity <- pnorm(ps_fit$train_predictions)For the remainder of this vignette we use a short synthetic run to keep build time low.
set.seed(1)
n <- 80; p <- 10
X_syn <- matrix(rnorm(n * p), n, p)
W_syn <- rbinom(n, 1, pnorm(0.8 * X_syn[, 1]))
ps_fit <- HorseTrees(
y = W_syn,
X_train = X_syn,
outcome_type = "binary",
number_of_trees = 5,
k = 0.5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
propensity_syn <- pnorm(ps_fit$train_predictions)
cat("Propensity scores — range: [",
round(range(propensity_syn), 3), "]\n")
#> Propensity scores — range: [ 0.402 0.551 ]HorseTrees() handles right-censored data via an AFT
model. Pass outcome_type = "right-censored" and provide the
status vector (1 = event, 0 = censored). When
timescale = "time" (the default), the package
log-transforms survival times internally and returns predictions on the
log scale (see Key Concepts above).
set.seed(2)
log_T <- X_syn[, 1] + rnorm(n)
C <- rexp(n, 0.5)
y_syn <- pmin(exp(log_T), C)
d_syn <- as.integer(exp(log_T) <= C)
ht_surv <- HorseTrees(
y = y_syn,
status = d_syn,
X_train = X_syn,
outcome_type = "right-censored",
timescale = "time",
number_of_trees = 5,
N_post = 50,
N_burn = 25,
store_posterior_sample = TRUE,
verbose = FALSE
)
cat("Posterior mean log-time (first 5 obs):",
round(ht_surv$train_predictions[1:5], 3), "\n")
#> Posterior mean log-time (first 5 obs): 1.245 1.403 1.489 1.654 1.279
cat("Posterior sigma — mean:",
round(mean(ht_surv$sigma), 3), "\n")
#> Posterior sigma — mean: 1.049When event times are not observed exactly but known to lie within an
interval, the package supports interval censoring.
Instead of providing y and status, pass
left_time and right_time with
outcome_type = "interval-censored".
The three censoring types are encoded as follows:
left_time[i] == right_time[i]left_time[i] < right_time[i] (both finite)right_time[i] = InfThis convention matches
survival::Surv(type = "interval2").
set.seed(20)
# Generate true event times
true_T <- rexp(n, rate = exp(-0.5 * X_syn[, 1]))
# Create interval-censored observations
left_syn <- true_T * runif(n, 0.5, 1.0)
right_syn <- true_T * runif(n, 1.0, 1.5)
# Mark some as exact observations and some as right-censored
exact_idx <- sample(n, 25)
left_syn[exact_idx] <- true_T[exact_idx]
right_syn[exact_idx] <- true_T[exact_idx]
rc_idx <- sample(setdiff(seq_len(n), exact_idx), 15)
right_syn[rc_idx] <- Inf
cat("Exact events:", sum(left_syn == right_syn), "\n")
#> Exact events: 25
cat("Interval-censored:", sum(left_syn < right_syn & is.finite(right_syn)), "\n")
#> Interval-censored: 40
cat("Right-censored:", sum(!is.finite(right_syn)), "\n")
#> Right-censored: 15
ht_ic <- HorseTrees(
left_time = left_syn,
right_time = right_syn,
X_train = X_syn,
outcome_type = "interval-censored",
timescale = "time",
number_of_trees = 5,
N_post = 50,
N_burn = 25,
store_posterior_sample = TRUE,
verbose = FALSE
)
cat("Posterior mean log-time (first 5 obs):",
round(ht_ic$train_predictions[1:5], 3), "\n")
#> Posterior mean log-time (first 5 obs): 0.79 0.839 0.769 0.989 0.799
cat("Posterior sigma — mean:",
round(mean(ht_ic$sigma), 3), "\n")
#> Posterior sigma — mean: 1.022All survival functions (SurvivalBART,
SurvivalDART, SurvivalBCF,
SurvivalShrinkageBCF) and the general-purpose functions
(ShrinkageTrees, CausalShrinkageForest,
CausalHorseForest) accept left_time and
right_time in the same way. For example, using
SurvivalBART:
While HorseTrees() fixes the prior to the horseshoe, the
more general ShrinkageTrees() function exposes the
prior_type argument, allowing the user to select among all
implemented regularisation strategies. Available options are
"horseshoe", "horseshoe_fw" (forest-wide),
"half-cauchy", "standard" (classical BART),
and "dirichlet" (DART). Below we compare the per-tree
horseshoe and the forest-wide horseshoe on a continuous outcome.
set.seed(3)
y_cont <- X_syn[, 1] + 0.5 * X_syn[, 2] + rnorm(n)
# Horseshoe prior (default for HorseTrees)
fit_hs <- ShrinkageTrees(
y = y_cont,
X_train = X_syn,
outcome_type = "continuous",
prior_type = "horseshoe",
local_hp = 1.0 / sqrt(5),
global_hp = 1.0 / sqrt(5),
number_of_trees = 5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
# Forest-wide horseshoe (horseshoe_fw)
fit_fw <- ShrinkageTrees(
y = y_cont,
X_train = X_syn,
outcome_type = "continuous",
prior_type = "horseshoe_fw",
local_hp = 1.0 / sqrt(5),
global_hp = 1.0 / sqrt(5),
number_of_trees = 5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
cat("Horseshoe — train RMSE:",
round(sqrt(mean((fit_hs$train_predictions - y_cont)^2)), 3), "\n")
#> Horseshoe — train RMSE: 1.201
cat("Horseshoe FW— train RMSE:",
round(sqrt(mean((fit_fw$train_predictions - y_cont)^2)), 3), "\n")
#> Horseshoe FW— train RMSE: 1.142SurvivalBART() and SurvivalDART() fit
classical BART and DART models for right-censored survival data under
the AFT formulation. They calibrate prior hyperparameters automatically
from the data range, providing a simple interface when horseshoe
shrinkage is not needed.
set.seed(4)
# SurvivalBART: classical BART prior, AFT likelihood
fit_sbart <- SurvivalBART(
time = y_syn,
status = d_syn,
X_train = X_syn,
number_of_trees = 5,
k = 2.0,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
# SurvivalDART: Dirichlet (DART) splitting prior
fit_sdart <- SurvivalDART(
time = y_syn,
status = d_syn,
X_train = X_syn,
number_of_trees = 5,
k = 2.0,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
cat("SurvivalBART class:", class(fit_sbart), "\n")
#> SurvivalBART class: ShrinkageTrees
cat("SurvivalDART class:", class(fit_sdart), "\n")
#> SurvivalDART class: ShrinkageTreesA key motivation for horseshoe shrinkage and the Dirichlet (DART) sparsity prior is their behaviour in the \(p \gg n\) regime: many covariates are available but only a small subset drives the outcome. Classical BART may struggle here because the standard Gaussian leaf prior is non-sparse and does not concentrate on a small number of predictors.
We illustrate both priors on a sparse AFT simulation: \(n = 60\) observations, \(p = 200\) predictors, and only three active predictors.
set.seed(20)
n_hd <- 60; p_hd <- 200
X_hd <- matrix(rnorm(n_hd * p_hd), n_hd, p_hd)
# True log-survival depends only on predictors 1, 2, and 3
log_T_hd <- 1.5 * X_hd[, 1] - 1.0 * X_hd[, 2] + 0.5 * X_hd[, 3] + rnorm(n_hd)
C_hd <- rexp(n_hd, rate = 0.5)
y_hd <- pmin(exp(log_T_hd), C_hd)
d_hd <- as.integer(exp(log_T_hd) <= C_hd)
cat("n =", n_hd, "| p =", p_hd,
"| active predictors = 3",
"| censoring rate =", round(1 - mean(d_hd), 2), "\n")
#> n = 60 | p = 200 | active predictors = 3 | censoring rate = 0.4ShrinkageTrees (horseshoe) places global–local shrinkage on the step heights of every leaf, automatically regularising all 200 predictors toward zero while preserving the signal in the three active ones.
set.seed(21)
fit_hd_hs <- ShrinkageTrees(
y = y_hd,
status = d_hd,
X_train = X_hd,
outcome_type = "right-censored",
prior_type = "horseshoe",
local_hp = 1.0 / sqrt(10),
global_hp = 1.0 / sqrt(10),
number_of_trees = 10,
N_post = 50,
N_burn = 25,
verbose = FALSE
)SurvivalDART uses a Dirichlet prior on split
probabilities to induce structural sparsity: after burn-in, most
splitting probability is concentrated on truly predictive variables.
Setting rho_dirichlet = 3 encodes the prior belief that
approximately three predictors are active.
set.seed(22)
fit_hd_dart <- SurvivalDART(
time = y_hd,
status = d_hd,
X_train = X_hd,
number_of_trees = 10,
rho_dirichlet = 3,
N_post = 50,
N_burn = 25,
verbose = FALSE
)Both models run without error in the \(p > n\) regime. We compare their posterior mean predictions in log-time against the latent true values used to generate the data.
rmse_hs <- sqrt(mean((fit_hd_hs$train_predictions - log_T_hd)^2))
rmse_dart <- sqrt(mean((fit_hd_dart$train_predictions - log_T_hd)^2))
cat(sprintf("%-18s train RMSE (log-time): %.3f\n", "Horseshoe", rmse_hs))
#> Horseshoe train RMSE (log-time): 2.531
cat(sprintf("%-18s train RMSE (log-time): %.3f\n", "DART", rmse_dart))
#> DART train RMSE (log-time): 2.402The DART model also produces variable importance plots that display the posterior distribution of each predictor’s splitting probability. With only 50 posterior draws the top-10 plot below should already concentrate most probability mass near the three truly active predictors.
For causal inference, ShrinkageTrees provides Bayesian Causal Forest (BCF) models that decompose the outcome into a prognostic component and a treatment effect component: \[ Y_i = \mu(\mathbf{x}_i) + W_i \cdot \tau(\mathbf{x}_i) + \varepsilon_i, \] where \(\mu(\cdot)\) is the prognostic (control) function modelled by one tree ensemble, and \(\tau(\cdot)\) is the heterogeneous treatment effect modelled by a second ensemble. This two-forest structure allows each component to have its own regularisation, number of trees, and prior — for instance, a standard BART prior for the prognostic forest and horseshoe shrinkage for the treatment effect forest.
The package provides four causal model functions with increasing
generality: SurvivalBCF() (classical BCF for survival),
SurvivalShrinkageBCF() (BCF + DART for survival),
CausalHorseForest() (horseshoe BCF for all outcome types),
and CausalShrinkageForest() (fully configurable BCF). We
illustrate each below on synthetic data with a known treatment
effect.
set.seed(5)
tau_true <- 1.5 * (X_syn[, 1] > 0) # heterogeneous treatment effect
y_causal <- X_syn[, 1] + W_syn * tau_true + rnorm(n, sd = 0.5)SurvivalBCF() fits a BCF model for right-censored
survival outcomes using classical BART priors.
# Full analysis (eval=FALSE — use larger MCMC settings in practice)
fit_sbcf <- SurvivalBCF(
time = time,
status = status,
X_train = X,
treatment = treatment,
propensity = propensity, # from HorseTrees above
N_post = 5000,
N_burn = 5000,
verbose = FALSE
)set.seed(6)
fit_sbcf <- SurvivalBCF(
time = y_syn,
status = d_syn,
X_train = X_syn,
treatment = W_syn,
number_of_trees_control = 5,
number_of_trees_treat = 5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
cat("SurvivalBCF class:", class(fit_sbcf), "\n")
#> SurvivalBCF class: CausalShrinkageForest
cat("ATE (posterior mean):",
round(mean(fit_sbcf$train_predictions_treat), 3), "\n")
#> ATE (posterior mean): 1.475SurvivalShrinkageBCF() extends BCF with a Dirichlet
splitting prior on both forests, inducing sparsity in high-dimensional
settings.
set.seed(7)
fit_ssbcf <- SurvivalShrinkageBCF(
time = y_syn,
status = d_syn,
X_train = X_syn,
treatment = W_syn,
number_of_trees_control = 5,
number_of_trees_treat = 5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
cat("SurvivalShrinkageBCF class:", class(fit_ssbcf), "\n")
#> SurvivalShrinkageBCF class: CausalShrinkageForestCausalHorseForest() is the primary novel contribution of
this package. It applies horseshoe shrinkage to the leaf parameters of
both the prognostic and treatment-effect forests. This enables effective
regularisation when many covariates are available but few are truly
predictive of heterogeneous treatment effects.
set.seed(8)
fit_chf <- CausalHorseForest(
y = y_causal,
X_train_control = X_syn,
X_train_treat = X_syn,
treatment_indicator_train = W_syn,
outcome_type = "continuous",
number_of_trees = 5,
N_post = 50,
N_burn = 25,
store_posterior_sample = TRUE,
verbose = FALSE
)
cat("CausalHorseForest class:", class(fit_chf), "\n")
#> CausalHorseForest class: CausalShrinkageForest
# Posterior mean CATE
cate_mean <- fit_chf$train_predictions_treat
cat("CATE — posterior mean (first 5):",
round(cate_mean[1:5], 3), "\n")
#> CATE — posterior mean (first 5): 1.581 1.507 1.48 1.538 1.466
# Posterior ATE
ate_samples <- rowMeans(fit_chf$train_predictions_sample_treat)
cat("ATE posterior mean:",
round(mean(ate_samples), 3),
" 95% CI: [",
round(quantile(ate_samples, 0.025), 3), ",",
round(quantile(ate_samples, 0.975), 3), "]\n")
#> ATE posterior mean: 1.505 95% CI: [ 0.894 , 2.134 ]The fitted object stores the posterior mean CATE for each training
observation in train_predictions_treat. When
store_posterior_sample = TRUE, the full posterior sample
matrix is available in train_predictions_sample_treat, from
which the posterior ATE distribution and credible intervals can be
computed as shown above.
You can also supply separate test matrices to obtain out-of-sample CATE predictions.
set.seed(9)
X_test <- matrix(rnorm(20 * p), 20, p)
W_test <- rbinom(20, 1, 0.5)
fit_chf_test <- CausalHorseForest(
y = y_causal,
X_train_control = X_syn,
X_train_treat = X_syn,
treatment_indicator_train = W_syn,
X_test_control = X_test,
X_test_treat = X_test,
treatment_indicator_test = W_test,
outcome_type = "continuous",
number_of_trees = 5,
N_post = 50,
N_burn = 25,
store_posterior_sample = TRUE,
verbose = FALSE
)
cat("Test CATE (first 5):",
round(fit_chf_test$test_predictions_treat[1:5], 3), "\n")
#> Test CATE (first 5): 1.378 1.269 1.553 1.425 1.416CausalShrinkageForest() is the most general causal model
interface. It allows independent prior choices for the prognostic and
treatment effect forests via prior_type_control and
prior_type_treat. For example, one could use a standard
BART prior for the prognostic forest (where variable selection is less
critical) and horseshoe shrinkage for the treatment forest (where most
covariates are expected to be irrelevant for the treatment effect).
set.seed(10)
lh <- 1.5 / sqrt(5)
fit_csf <- CausalShrinkageForest(
y = y_causal,
X_train_control = X_syn,
X_train_treat = X_syn,
treatment_indicator_train = W_syn,
outcome_type = "continuous",
prior_type_control = "horseshoe",
prior_type_treat = "horseshoe",
local_hp_control = lh,
global_hp_control = lh,
local_hp_treat = lh,
global_hp_treat = lh,
number_of_trees_control = 5,
number_of_trees_treat = 5,
N_post = 50,
N_burn = 25,
store_posterior_sample = TRUE,
verbose = FALSE
)
cat("CausalShrinkageForest class:", class(fit_csf), "\n")
#> CausalShrinkageForest class: CausalShrinkageForest
cat("Acceptance ratio (control):",
round(fit_csf$acceptance_ratio_control, 3), "\n")
#> Acceptance ratio (control): 0.368
cat("Acceptance ratio (treat) :",
round(fit_csf$acceptance_ratio_treat, 3), "\n")
#> Acceptance ratio (treat) : 0.288The horseshoe_fw prior adds a forest-wide shrinkage
parameter that is tracked in the fitted object.
set.seed(11)
fit_fw2 <- CausalShrinkageForest(
y = y_causal,
X_train_control = X_syn,
X_train_treat = X_syn,
treatment_indicator_train = W_syn,
outcome_type = "continuous",
prior_type_control = "horseshoe_fw",
prior_type_treat = "horseshoe_fw",
local_hp_control = lh,
global_hp_control = lh,
local_hp_treat = lh,
global_hp_treat = lh,
number_of_trees_control = 5,
number_of_trees_treat = 5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)
cat("Forest-wide shrinkage (control, first 5 draws):\n")
#> Forest-wide shrinkage (control, first 5 draws):
print(round(fit_fw2$forestwide_shrinkage_control[1:5], 4))
#> [1] 1 1 1 1 1A fitted tree ensemble is hard to interpret. The posterior over \(f(x)\) does not say which covariates drive
the fit, or how. We implement the posterior summarisation of Woody et
al. (2021). The method projects every posterior draw of the fitted
function onto a simpler model, \[
\gamma^{(s)} = \arg\min_{\gamma \in \mathcal{G}}
\sum_i w_i \left[ f^{(s)}(x_i) - \gamma(x_i) \right]^2,
\qquad s = 1, \ldots, S.
\] The result is a posterior distribution over simple summaries:
coefficients with credible intervals, plus the posterior of the summary
\(R^2\). The summary \(R^2\) reports how much of the ensemble the
summary captures. We do not refit the model. The projection reads the
stored draws (store_posterior_sample = TRUE), so it takes
seconds.
The penalised families are the exception. We project the posterior mean of \(f\) alone there, and we report no intervals, for the reasons given in the section on penalised projections below.
We make two design choices. First, the function does no covariate
selection. By default it uses every covariate, and you choose a subset
with covariates. The coefficient table then describes
exactly the model you asked for. Second, there is no
newdata argument. We project the fitted object over the
design it was fitted to.
We illustrate on the high-dimensional survival fit from the previous section (\(n = 60\), \(p = 200\), three active predictors).
set.seed(30)
fit_pp <- HorseTrees(
y = y_hd,
status = d_hd,
X_train = X_hd,
outcome_type = "right-censored",
timescale = "time",
number_of_trees = 5,
N_post = 50,
N_burn = 25,
store_posterior_sample = TRUE,
verbose = FALSE
)The default family is a linear model. The output follows
summary(lm(...)). All uncertainty is posterior: we report
credible intervals, not test statistics or p-values.
pp_lin <- posterior_projection(fit_pp, family = "linear",
covariates = 1:5)
pp_lin
#>
#> Projection residuals:
#> Min 1Q Median 3Q Max
#> -0.15200 -0.05407 -0.01025 0.05721 0.27400
#>
#> Coefficients (log-time (inverted from stored time scale)):
#> Estimate Post.SD 2.5% 97.5%
#> (Intercept) -0.510200 0.30170 -1.11600 0.01697
#> X1 0.047090 0.10800 -0.08589 0.35980
#> X2 0.007744 0.05369 -0.07167 0.12460
#> X3 0.005763 0.06503 -0.15410 0.15370
#> X4 -0.007766 0.05092 -0.12820 0.07503
#> X5 0.030560 0.08198 -0.08027 0.19050
#>
#> Residual standard error: 0.08128 on 54 degrees of freedom
#> Summary R-squared: 0.1398 (95% CI: 0.03379, 0.3182)
#> Draws: 50 Observations: 60 Covariates: 5The summary R-squared is the headline number. It gives the share of variation in the fitted function that this five-covariate summary captures. The credible interval shows how consistent that share is across the posterior.
We store the coefficient draws, so any posterior quantity is available directly.
quantile(pp_lin$coefficients[, 2], c(0.025, 0.5, 0.975))
#> 2.5% 50% 97.5%
#> -0.08589289 0.01156105 0.35980868With \(p \ge n\) the unpenalised projection onto all covariates is undefined. The function stops with a message rather than return a meaningless answer.
posterior_projection(fit_pp, family = "linear")
#> Error:
#> ! The projection needs 201 basis columns but there are only 60 observations, so it is undefined.
#> Either pass `covariates = ` a subset of at most 58 basis columns' worth,
#> or use penalty = "ridge" to project onto all 200 covariates with regularisation, which summarises the posterior mean
#> and reports no intervals.The penalised options keep the projection defined when \(p \ge n\). We solve the penalised problem
once, on the posterior mean of the fitted surface, and we report point
estimates only. This is a deliberate restriction. A penalised estimator
is biased and non-smooth in its input, so the spread of penalised
solutions across draws has no coverage guarantee for the projection of
the true surface, and the fraction of draws selecting a covariate reads
as a selection probability while being an artefact of the penalty. We
prefer a single sparse summary that states what it is. Use
penalty = "none" when the uncertainty in the summary is the
point.
We choose \(\lambda\) by
cv.glmnet on the posterior mean, at
lambda.1se. Set lambda = "min" or a number to
override this.
pp_las <- posterior_projection(fit_pp, family = "linear",
penalty = "lasso")
pp_las
#>
#> Projection residuals:
#> Min 1Q Median 3Q Max
#> -0.13660 -0.06593 -0.02632 0.05604 0.36410
#>
#> Measure: Mean-Squared Error
#> Lambda Index Measure SE Nonzero
#> min 0.023227 19 0.007520 0.001882 4
#> 1se 0.044548 5 0.009379 0.002357 2
#> Using lambda.1se
#>
#> Coefficients (log-time (inverted from stored time scale)):
#> Estimate
#> (Intercept) -0.508000
#> X1 0.005360
#> X143 0.007551
#>
#> Residual standard error: 0.09333 on 57 degrees of freedom
#> Summary R-squared: 0.1303
#> Projection of the posterior mean over 50 draws (no uncertainty quantification)
#> Observations: 60 Covariates: 200
# One set of coefficients, most of them exactly zero:
sum(pp_las$coefficients[1, -1] != 0)
#> [1] 2penalty = "ridge" keeps all covariates. It is a
regularised projection, not a lower-dimensional one.
penalty = "elastic_net" sits between the two; set the mix
with alpha.
family = "additive" projects onto a sum of per-covariate
functions. We build each one from a natural spline basis. This is our
substitute for a partial dependence plot: one effect curve per
covariate, with posterior credible bands, from the stored draws. Note
that it is a projection, not a partial dependence integral.
pp_add <- posterior_projection(fit_pp, family = "additive",
covariates = 1:3, df_spline = 3)
pp_add
#>
#> Projection residuals:
#> Min 1Q Median 3Q Max
#> -0.132100 -0.056280 -0.007884 0.044350 0.224600
#>
#> Partial effects (log-time (inverted from stored time scale)), 10th to 90th percentile of each covariate:
#> Estimate Post.SD 2.5% 97.5%
#> X1 0.14700 0.3268 -0.1683 1.1010
#> X2 0.02213 0.1452 -0.2436 0.3013
#> X3 -0.02702 0.2101 -0.6531 0.2100
#> Basis coefficients are in x$coef_summary; plot(type = "effects") draws the curves.
#>
#> Residual standard error: 0.07915 on 50 degrees of freedom
#> Summary R-squared: 0.1802 (95% CI: 0.03319, 0.3535)
#> Draws: 50 Observations: 60 Covariates: 3The printed table gives, for each covariate, the posterior of how far its effect moves between the 10th and 90th percentile of that covariate. This is a nonlinear analogue of a slope. We plot the curves themselves with:
family = "tree" projects onto a shallow CART partition.
The leaves are interpretable subgroups. Each leaf has a posterior
credible interval for its average fitted value.
pp_tree <- posterior_projection(fit_pp, family = "tree", max_leaves = 4)
pp_tree
#>
#> Projection residuals:
#> Min 1Q Median 3Q Max
#> -0.098600 -0.033930 -0.004704 0.022140 0.260600
#>
#> Coefficients (log-time (inverted from stored time scale)):
#> Estimate Post.SD 2.5% 97.5%
#> leaf1 -0.5811 0.2905 -1.143 -0.082860
#> leaf2 -0.5348 0.2747 -1.086 0.001394
#> leaf3 -0.3914 0.6194 -1.038 1.724000
#> leaf4 -0.3878 0.5343 -1.113 1.079000
#>
#> Residual standard error: 0.055 on 56 degrees of freedom
#> Summary R-squared: 0.1833 (95% CI: 0.005465, 0.9609)
#> Draws: 50 Observations: 60 Covariates: 200
#> Leaves: 4 Split on: X143, X88, X136For a CausalShrinkageForest, target selects
the surface. "treatment" projects the CATE surface \(\tau(x)\). This answers which covariates
the treatment effect varies along, a question the ATE and CATE SD from
summary() cannot address. "prognostic"
projects \(\mu(x)\).
pp_tau <- posterior_projection(fit_chf, target = "treatment",
family = "linear", covariates = 1:5)
pp_tau
#>
#> Projection residuals:
#> Min 1Q Median 3Q Max
#> -0.28970 -0.07962 0.01021 0.07700 0.22400
#>
#> Coefficients (response):
#> Estimate Post.SD 2.5% 97.5%
#> (Intercept) 1.501000 0.38350 0.8994 2.1400
#> X1 0.018700 0.09592 -0.1494 0.2602
#> X2 -0.039330 0.13560 -0.2530 0.2203
#> X3 0.004848 0.10350 -0.1989 0.1745
#> X4 -0.025110 0.12830 -0.2850 0.2953
#> X5 0.019860 0.16970 -0.2105 0.3588
#>
#> Residual standard error: 0.1126 on 74 degrees of freedom
#> Summary R-squared: 0.3139 (95% CI: 0.06174, 0.6276)
#> Draws: 50 Observations: 80 Covariates: 5# Subgroups of treatment effect, with posterior intervals per leaf
posterior_projection(fit_chf, target = "treatment",
family = "tree", max_leaves = 4)
#>
#> Projection residuals:
#> Min 1Q Median 3Q Max
#> -0.199500 -0.049920 -0.002167 0.056360 0.180300
#>
#> Coefficients (response):
#> Estimate Post.SD 2.5% 97.5%
#> leaf1 1.270 0.5399 0.2804 2.123
#> leaf2 1.419 0.4756 0.5288 2.135
#> leaf3 1.508 0.4294 0.8621 2.351
#> leaf4 1.610 0.4620 0.8223 2.537
#>
#> Residual standard error: 0.07941 on 76 degrees of freedom
#> Summary R-squared: 0.225 (95% CI: 0.01509, 0.844)
#> Draws: 50 Observations: 80 Covariates: 10
#> Leaves: 4 Split on: X7, X5, X8We handle the scale automatically. With
timescale = "time" the stored draws are on the time scale,
so we take logs before we project. The output labels the coefficients as
contributions to the log acceleration factor.
posterior_projection() is a generic. There is also a
default method, which takes a plain S by n
matrix of posterior draws together with the design they correspond to.
Use it to project draws from another source, or from a part of a fitted
object that the two methods above do not reach. The result is an object
of class PosteriorProjection, with its own
print() and plot() methods.
All fitted objects — whether from prediction models or causal models
— support a consistent set of S3 methods: print(),
summary(), predict(), and plot().
This section illustrates each method using the models fitted above. The
package adds one further generic, posterior_projection(),
covered in the previous section.
Calling print() (or just typing the object name)
displays a concise model summary.
print(fit_chf)
#>
#> CausalShrinkageForest model
#> ---------------------------
#> Outcome type: Continuous
#> Training size (n): 80
#> Posterior draws: 50 (burn-in 25)
#> Posterior mean sigma: 0.802
#>
#> Control Treatment
#> ------------------- -------------------
#> Prior: horseshoe horseshoe
#> Number of trees: 5 5
#> Number of features: 10 10
#> Acceptance ratio: 0.46 0.412For causal models the output additionally shows the number of trees in each forest and prior details for both components.
print(fit_csf)
#>
#> CausalShrinkageForest model
#> ---------------------------
#> Outcome type: Continuous
#> Training size (n): 80
#> Posterior draws: 50 (burn-in 25)
#> Posterior mean sigma: 0.817
#>
#> Control Treatment
#> ------------------- -------------------
#> Prior: horseshoe horseshoe
#> Number of trees: 5 5
#> Number of features: 10 10
#> Acceptance ratio: 0.368 0.288summary() returns a structured list and displays a
richer description including posterior statistics for \(\sigma\), acceptance ratios, and treatment
effect estimates.
smry <- summary(fit_hs)
print(smry)
#>
#> ShrinkageTrees model summary
#> ============================
#> Call: ShrinkageTrees(y = y_cont, X_train = X_syn, outcome_type = "continuous",
#> number_of_trees = 5, prior_type = "horseshoe", local_hp = 1/sqrt(5),
#> global_hp = 1/sqrt(5), N_post = 50, N_burn = 25, verbose = FALSE)
#>
#> Outcome: Continuous | Prior: horseshoe | Trees: 5
#> Data: n = 80, p = 10 | Draws: 50 (burn-in 25)
#>
#> Posterior sigma:
#> Mean: 0.984 SD: 0.099 95% CI: [0.816, 1.218]
#>
#> Predictions (posterior mean):
#> Train: mean = 0.052, sd = 0.098, range = [-0.298, 0.212]
#> Test: mean = 0.067, sd = NA, range = [0.067, 0.067]
#>
#> Variable importance (posterior inclusion probability):
#> X8: 0.205 X4: 0.173 X5: 0.134 X9: 0.126 X7: 0.124 X10: 0.078 X3: 0.077 X2: 0.06 X6: 0.02 X1: 0.003
#>
#> MCMC acceptance ratio: 0.504
#>
#> Convergence diagnostics (coda):
#> Effective sample size: sigma = 50For causal models the summary includes the posterior ATE with a 95%
credible interval (when store_posterior_sample = TRUE).
smry_c <- summary(fit_chf)
print(smry_c)
#>
#> CausalShrinkageForest model summary
#> =====================================
#> Call: CausalHorseForest(y = y_causal, X_train_control = X_syn, X_train_treat = X_syn,
#> treatment_indicator_train = W_syn, outcome_type = "continuous",
#> number_of_trees = 5, N_post = 50, N_burn = 25, store_posterior_sample = TRUE,
#> verbose = FALSE)
#>
#> Outcome: Continuous
#> Prior: control = horseshoe, treatment = horseshoe
#> Trees: control = 5, treatment = 5
#> Data: n = 80, p_control = 10, p_treat = 10 | Draws: 50 (burn-in 25)
#>
#> Treatment effect:
#> PATE: 1.4909 95% CI (Bayesian bootstrap): [0.7988, 2.1264]
#> CATE SD: 0.1194
#>
#> Prognostic function (mu):
#> Mean: 0.675 SD: 0.137 Range: [0.419, 0.895]
#>
#> Posterior sigma:
#> Mean: 0.802 SD: 0.089 95% CI: [0.633, 0.989]
#>
#> Variable importance - control forest (posterior inclusion probability):
#> X8: 0.18 X6: 0.155 X5: 0.128 X4: 0.125 X9: 0.089 X2: 0.077 X7: 0.068 X10: 0.068 X1: 0.06 X3: 0.051
#>
#> Variable importance - treatment forest (posterior inclusion probability):
#> X5: 0.151 X8: 0.141 X9: 0.124 X3: 0.111 X1: 0.107 X10: 0.094 X7: 0.091 X2: 0.065 X4: 0.062 X6: 0.054
#>
#> MCMC acceptance ratios: control = 0.46, treatment = 0.412
# Access the ATE directly
cat("ATE mean :", round(smry_c$treatment_effect$ate, 3), "\n")
#> ATE mean : 1.491
cat("ATE 95% CI: [",
round(smry_c$treatment_effect$ate_lower, 3), ",",
round(smry_c$treatment_effect$ate_upper, 3), "]\n")
#> ATE 95% CI: [ 0.799 , 2.126 ]By default the ATE credible interval is obtained by a Bayesian bootstrap: at each MCMC iteration \(s\) the observation-level CATEs \(\tau^{(s)}(x_i)\) are reweighted with Dirichlet(1, …, 1) weights,
\[ \widehat{\mathrm{PATE}}^{(s)} \;=\; \sum_{i=1}^n w_i^{(s)}\, \tau^{(s)}(x_i), \qquad (w_1^{(s)}, \dots, w_n^{(s)}) \sim \mathrm{Dir}(1, \dots, 1). \]
The collection \(\{\widehat{\mathrm{PATE}}^{(s)}\}\)
approximates the posterior of the population ATE and
therefore propagates uncertainty in both \(\tau(\cdot)\) and the covariate
distribution \(F_X\). Setting
bayesian_bootstrap = FALSE reverts to equal \(1/n\) weights, giving the mixed
ATE (MATE) that conditions on the observed covariates and has a
narrower credible interval.
smry_pate <- summary(fit_chf, bayesian_bootstrap = TRUE) # default
smry_mate <- summary(fit_chf, bayesian_bootstrap = FALSE)The standalone helper bayesian_bootstrap_ate() returns
both posteriors and their draws in a single list, and also works on a
CausalShrinkageForestPrediction returned by
predict() so that the PATE integrates over a prespecified
target population.
bb <- bayesian_bootstrap_ate(fit_chf)
cat("PATE:", round(bb$pate_mean, 3),
" 95% CI: [", round(bb$pate_ci$lower, 3), ",",
round(bb$pate_ci$upper, 3), "]\n")
#> PATE: 1.514 95% CI: [ 0.875 , 2.15 ]
cat("MATE:", round(bb$mate_mean, 3),
" 95% CI: [", round(bb$mate_ci$lower, 3), ",",
round(bb$mate_ci$upper, 3), "]\n")
#> MATE: 1.505 95% CI: [ 0.894 , 2.134 ]predict() computes the posterior predictive distribution
on new data. It returns a ShrinkageTreesPrediction object
with posterior mean and credible-interval vectors.
X_new <- matrix(rnorm(10 * p), 10, p)
pred <- predict(fit_hs, newdata = X_new)
print(pred)
#>
#> ShrinkageTrees predictions
#> --------------------------
#> Observations: 10
#> Credible interval: 95%
#> Scale: fitted value
#>
#> mean lower upper
#> -------- -------- --------
#> [ 1] 0.105 -0.423 0.740
#> [ 2] 0.053 -0.503 0.618
#> [ 3] -0.128 -0.736 0.292
#> [ 4] 0.058 -0.369 0.528
#> [ 5] 0.015 -0.424 0.487
#> [ 6] -0.007 -0.464 0.485
#> ... (4 more)# Point estimates and 95% credible intervals
head(data.frame(
mean = round(pred$mean, 3),
lower = round(pred$lower, 3),
upper = round(pred$upper, 3)
))
#> mean lower upper
#> 1 0.105 -0.423 0.740
#> 2 0.053 -0.503 0.618
#> 3 -0.128 -0.736 0.292
#> 4 0.058 -0.369 0.528
#> 5 0.015 -0.424 0.487
#> 6 -0.007 -0.464 0.485For causal models (CausalShrinkageForest and
CausalHorseForest), predict() returns three
sets of posterior summaries:
For survival models with timescale = "time", the
prognostic and total components are back-transformed to the original
time scale (via \(\exp(\cdot)\)), and
the CATE becomes a multiplicative time ratio: \(\exp(\tau) > 1\) means treatment
prolongs survival.
The predict() method requires two covariate matrices —
one for each forest — matching the columns used at fit time.
X_new_ctrl <- matrix(rnorm(10 * p), 10, p)
X_new_treat <- matrix(rnorm(10 * p), 10, p)
pred_c <- predict(fit_chf, newdata_control = X_new_ctrl,
newdata_treat = X_new_treat)
print(pred_c)
#>
#> CausalShrinkageForest predictions
#> ----------------------------------
#> Observations: 10
#> Credible interval: 95%
#> Outcome type: continuous
#>
#> PATE: 1.565 95% CI (Bayesian bootstrap): [0.796, 2.396]
#>
#> Prognostic (mu):
#> mean lower upper
#> -------- -------- --------
#> [ 1] 0.718 -0.048 1.415
#> [ 2] 0.795 0.216 1.611
#> [ 3] 0.726 -0.201 1.469
#> [ 4] 0.475 -0.661 1.184
#> [ 5] 0.609 -0.645 1.667
#> [ 6] 0.597 -0.621 1.342
#> ... (4 more)
#>
#> CATE (tau):
#> mean lower upper
#> -------- -------- --------
#> [ 1] 1.526 0.390 2.534
#> [ 2] 1.739 0.802 2.658
#> [ 3] 1.148 -0.200 2.455
#> [ 4] 1.653 0.776 2.615
#> [ 5] 1.647 0.688 2.537
#> [ 6] 1.498 -0.052 2.883
#> ... (4 more)
#>
#> Total (mu + tau):
#> mean lower upper
#> -------- -------- --------
#> [ 1] 2.244 0.667 3.397
#> [ 2] 2.534 1.448 3.653
#> [ 3] 1.874 0.416 3.297
#> [ 4] 2.128 0.863 3.576
#> [ 5] 2.256 0.682 4.112
#> [ 6] 2.094 0.502 3.637
#> ... (4 more)# Extract individual components
head(data.frame(
prognostic = round(pred_c$prognostic$mean, 3),
cate = round(pred_c$cate$mean, 3),
total = round(pred_c$total$mean, 3)
))
#> prognostic cate total
#> 1 0.718 1.526 2.244
#> 2 0.795 1.739 2.534
#> 3 0.726 1.148 1.874
#> 4 0.475 1.653 2.128
#> 5 0.609 1.647 2.256
#> 6 0.597 1.498 2.094The plot() method produces diagnostic and inferential
graphics using the ggplot2 package (a suggested
dependency).
The ATE density uses the Bayesian-bootstrap PATE posterior by
default; pass bayesian_bootstrap = FALSE to plot the
(narrower) mixed ATE density instead. See the summary section above for
the definitions.
Variable importance plots are available when
prior_type = "dirichlet". For causal models with
prior_type_control = "dirichlet" or
prior_type_treat = "dirichlet", use
forest = "control", "treat", or
"both".
set.seed(12)
fit_dart <- ShrinkageTrees(
y = y_cont,
X_train = X_syn,
outcome_type = "continuous",
prior_type = "dirichlet",
local_hp = 1.0 / sqrt(5),
number_of_trees = 5,
N_post = 50,
N_burn = 25,
verbose = FALSE
)For survival models (outcome_type = "right-censored" or
"interval-censored"), the plot() method can
draw posterior survival curves derived from the fitted AFT log-normal
model: \[
S(t \mid \mathbf{x}_i) = 1 - \Phi\!\left(\frac{\log t -
\mu_i}{\sigma}\right),
\] where \(\mu_i =
f(\mathbf{x}_i)\) is the BART ensemble prediction and \(\sigma\) is the residual standard deviation
on the log-time scale.
The type = "survival" option supports two modes
controlled by the obs argument:
obs = NULL,
the default): computes \(\bar{S}(t) =
n^{-1}\sum_i S(t \mid \mathbf{x}_i)\) at each MCMC iteration,
giving credible bands that reflect full posterior uncertainty.obs = c(1, 5, ...)): one curve per selected training
observation with its own credible band.Additional options:
| Argument | Description |
|---|---|
level |
Width of the pointwise credible band (default
0.95). |
t_grid |
Custom time grid (original scale). Auto-generated if
NULL. |
km |
If TRUE, overlay the Kaplan–Meier estimate
(population-average only). |
We use the survival fit from the earlier section:
# Same curve with the Kaplan-Meier estimate overlaid for comparison
plot(ht_surv, type = "survival", km = TRUE)# Individual survival curves for observations 1, 20, 40, 60, and 80
plot(ht_surv, type = "survival", obs = c(1, 20, 40, 60, 80))# Single individual with a narrower 90% credible band
plot(ht_surv, type = "survival", obs = 1, level = 0.90)When store_posterior_sample = FALSE, the credible bands
only reflect uncertainty in \(\sigma\)
(using plug-in posterior mean \(\hat{\mu}_i\)). The survival functions
(SurvivalBART, SurvivalDART, etc.) store
posterior samples by default, so full posterior bands are available out
of the box.
The survival curves above are based on the training
data — they show \(S(t \mid
\mathbf{x}_i)\) for the observations used to fit the model. For
new (out-of-sample) data, call predict()
first and then plot() on the prediction object. This
produces posterior predictive survival curves that propagate
full parameter uncertainty through to the new covariate values:
# New observations for prediction
set.seed(99)
X_new <- matrix(rnorm(20 * p), ncol = p)
pred_surv <- predict(ht_surv, newdata = X_new)# Individual posterior predictive curves for selected new observations
plot(pred_surv, type = "survival", obs = c(1, 5, 10))The same level and t_grid arguments are
available as for the training-data survival curves. The Kaplan–Meier
overlay (km = TRUE) is not available for prediction
objects, since observed event times are only known for the training
set.
Running multiple independent chains improves mixing diagnostics and
reduces sensitivity to starting values. Pass
n_chains > 1 to any model-fitting function; chains are
run in parallel via parallel::mclapply on Unix-like
systems.
set.seed(13)
fit_2chain <- ShrinkageTrees(
y = y_cont,
X_train = X_syn,
outcome_type = "continuous",
prior_type = "horseshoe",
local_hp = 1.0 / sqrt(5),
global_hp = 1.0 / sqrt(5),
number_of_trees = 5,
N_post = 50,
N_burn = 25,
n_chains = 2,
verbose = FALSE
)
cat("n_chains stored :", fit_2chain$mcmc$n_chains, "\n")
#> n_chains stored : 2
cat("Total sigma draws:", length(fit_2chain$sigma),
" (2 chains x 50 draws)\n")
#> Total sigma draws: 100 (2 chains x 50 draws)
cat("Per-chain acceptance ratios:\n")
#> Per-chain acceptance ratios:
print(round(fit_2chain$chains$acceptance_ratios, 3))
#> [1] 0.532 0.388The same interface works for causal models.
set.seed(14)
fit_causal_2chain <- CausalShrinkageForest(
y = y_causal,
X_train_control = X_syn,
X_train_treat = X_syn,
treatment_indicator_train = W_syn,
outcome_type = "continuous",
prior_type_control = "horseshoe",
prior_type_treat = "horseshoe",
local_hp_control = lh,
global_hp_control = lh,
local_hp_treat = lh,
global_hp_treat = lh,
number_of_trees_control = 5,
number_of_trees_treat = 5,
N_post = 50,
N_burn = 25,
n_chains = 2,
verbose = FALSE
)
cat("Pooled sigma draws:", length(fit_causal_2chain$sigma), "\n")
#> Pooled sigma draws: 100
cat("Per-chain acceptance ratios (control):\n")
#> Per-chain acceptance ratios (control):
print(round(fit_causal_2chain$chains$acceptance_ratios_control, 3))
#> [1] 0.328 0.328With multiple chains the traceplot shows one line per chain, and the
overlaid density plot compares the marginal posterior of \(\sigma\) across chains — both require
n_chains > 1.
MCMC methods require careful assessment of convergence before trusting the posterior summaries. Here are practical guidelines for ShrinkageTrees models.
The traceplot of the error standard deviation \(\sigma\) (via
plot(fit, type = "trace")) is the primary diagnostic. A
well-mixing chain should show:
N_burn.n_chains > 1): separate chains should overlap
substantially. The density overlay
(plot(fit, type = "density")) makes this easy to check
visually.The summary() output reports the average
Metropolis–Hastings acceptance ratio for the tree structure proposals
(grow/prune moves). As a rough guide:
N_burn and N_post, or relaxing the
tree structure prior (lower power, higher
base).When the suggested package coda is installed,
summary() automatically reports effective sample
size (ESS) and — for multi-chain fits — the
Gelman–Rubin \(\hat{R}\).
# summary() includes convergence diagnostics when coda is available
summary(fit_2chain)
#>
#> ShrinkageTrees model summary
#> ============================
#> Call: ShrinkageTrees(y = y_cont, X_train = X_syn, outcome_type = "continuous",
#> number_of_trees = 5, prior_type = "horseshoe", local_hp = 1/sqrt(5),
#> global_hp = 1/sqrt(5), N_post = 50, N_burn = 25, n_chains = 2,
#> verbose = FALSE)
#>
#> Outcome: Continuous | Prior: horseshoe | Trees: 5
#> Data: n = 80, p = 10 | Draws: 50 x 2 chains (burn-in 25)
#>
#> Posterior sigma:
#> Mean: 0.972 SD: 0.091 95% CI: [0.798, 1.159]
#>
#> Predictions (posterior mean):
#> Train: mean = 0.039, sd = 0.096, range = [-0.276, 0.242]
#> Test: mean = 0.053, sd = NA, range = [0.053, 0.053]
#>
#> Variable importance (posterior inclusion probability):
#> X3: 0.149 X2: 0.129 X5: 0.117 X8: 0.113 X4: 0.108 X7: 0.086 X10: 0.079 X6: 0.079 X1: 0.075 X9: 0.066
#>
#> MCMC acceptance ratio (per chain): 0.532, 0.388
#>
#> Convergence diagnostics (coda):
#> Gelman-Rubin R-hat: = 1.029
#> Effective sample size: sigma = 73For more detailed diagnostics, convert the fitted object to a
coda::mcmc.list with as.mcmc.list():
library(coda)
mcmc_obj <- as.mcmc.list(fit_2chain)
# Gelman-Rubin R-hat (values near 1 indicate convergence)
coda::gelman.diag(mcmc_obj)
#> Potential scale reduction factors:
#>
#> Point est. Upper C.I.
#> sigma 1.01 1.06
# Effective sample size
coda::effectiveSize(mcmc_obj)
#> sigma
#> 73.48486
# Geweke diagnostic (per chain)
coda::geweke.diag(mcmc_obj[[1]])
#>
#> Fraction in 1st window = 0.1
#> Fraction in 2nd window = 0.5
#>
#> sigma
#> -1.889The returned mcmc.list object is compatible with all
coda functions, including
coda::autocorr.plot(), coda::gelman.plot(),
coda::heidel.diag(), and
coda::raftery.diag().
The examples in this vignette use very small N_post and
N_burn to keep build time low. For a real analysis:
N_post = 2000, N_burn = 2000.N_post = 5000, N_burn = 5000.N_post = 5000, N_burn = 10000 (the AFT data augmentation
step can slow mixing, so a longer burn-in helps).n_chains = 2 or
4 to verify convergence and produce pooled posterior
samples.The full analysis of the pdac dataset replicates the
case study from Jacobs, van Wieringen & van der Pas (2026). Due to
the high-dimensional covariate space (~3,000 genes) and the large MCMC
settings needed for reliable inference, the code below is provided for
reference but is not evaluated during vignette building. Pre-computed
results can be reproduced by running the pdac_analysis
demo:
demo("pdac_analysis", package = "ShrinkageTrees").
data("pdac")
time <- pdac$time
status <- pdac$status
treatment <- pdac$treatment
X <- as.matrix(pdac[, !(names(pdac) %in% c("time","status","treatment"))])
set.seed(2025)
ps_fit <- HorseTrees(
y = treatment,
X_train = X,
outcome_type = "binary",
k = 1.0,
N_post = 5000,
N_burn = 5000,
verbose = FALSE
)
propensity <- pnorm(ps_fit$train_predictions)# Overlap plot
p0 <- propensity[treatment == 0]
p1 <- propensity[treatment == 1]
hist(p0, breaks = 15, col = rgb(1, 0.5, 0, 0.5), xlim = range(propensity),
xlab = "Propensity score", main = "Propensity score overlap")
hist(p1, breaks = 15, col = rgb(0, 0.5, 0, 0.5), add = TRUE)
legend("topright", legend = c("Control", "Treated"),
fill = c(rgb(1,0.5,0,0.5), rgb(0,0.5,0,0.5)))# Augment control matrix with propensity scores (BCF-style)
X_control <- cbind(propensity, X)
# Log-transform and centre survival times
log_time <- log(time) - mean(log(time))
set.seed(2025)
fit_pdac <- CausalHorseForest(
y = log_time,
status = status,
X_train_control = X_control,
X_train_treat = X,
treatment_indicator_train = treatment,
outcome_type = "right-censored",
timescale = "log",
number_of_trees = 200,
N_post = 5000,
N_burn = 5000,
store_posterior_sample = TRUE,
verbose = FALSE
)Chipman, H. A., George, E. I., & McCulloch, R. E. (2010). Bayesian Additive Regression Trees. Annals of Applied Statistics, 4(1), 266–298.
Hahn, P. R., Murray, J. S., & Carvalho, C. M. (2020). Bayesian Regression Tree Models for Causal Inference: Regularization, Confounding, and Heterogeneous Treatment Effects. Bayesian Analysis, 15(3), 965–1056.
Jacobs, T., van Wieringen, W. N., & van der Pas, S. L. (2026). Horseshoe Forests for High-Dimensional Causal Survival Analysis. Bayesian Analysis, advance publication, 1–30.
Linero, A. R. (2018). Bayesian Regression Trees for High-Dimensional Prediction and Variable Selection. Journal of the American Statistical Association, 113(522), 626–636.
Sparapani, R., Spanbauer, C., & McCulloch, R. (2021). Nonparametric Machine Learning and Efficient Computation with Bayesian Additive Regression Trees: The BART R Package. Journal of Statistical Software, 97(1), 1–66.
Woody, S., Carvalho, C. M., & Murray, J. S. (2021). Model Interpretation Through Lower-Dimensional Posterior Summarization. Journal of Computational and Graphical Statistics, 30(1), 144-161.