models.HierarchicalStacking

Hierarchical stacking (Yao, Pirš, Vehtari and Gelman, 2022): model

Usage

models.HierarchicalStacking()

weights that vary with covariates, w = softmax(alpha + B x) against the last model as reference, so a model can be trusted in one part of the portfolio and not another. The priors are those of BayesBlend’s HierarchicalBayesStacking; sampled by NUTS.

Scale continuous covariates (BayesBlend divides by twice the standard deviation) and dummy-code discrete ones before fitting, with the dummies first.

With partial_pooling, each model’s slopes on the discrete covariates, and separately on the continuous ones, are drawn around a model-level mean, itself drawn around a global mean. A scale of 0 removes a level: tau_mu_global=0 fixes the global mean at 0, tau_mu_*=0 pools completely and tau_sigma_*=0 sets every slope to its model’s mean. BayesBlend warns that pooling needs at least three covariates. adaptive multiplies the prior scales by N**lambda with lambda ~ Exponential(adaptive), weakening them as the data grow.

Parameters

discrete: int = 0

Number of leading covariates that are dummy codes.

alpha_loc: float = 0.0, 1.0
alpha_scale: float = 0.0, 1.0
beta_loc: float = 0.0, 1.0

Slope prior without pooling.

beta_scale: float = 0.0, 1.0

Slope prior without pooling.

partial_pooling: bool = False
tau_mu_global: float = 1.0
tau_mu_discrete: float = 1.0
tau_mu_continuous: float = 1.0
tau_sigma_discrete: float = 1.0

Pooling scales, BayesBlend’s defaults.

tau_sigma_continuous: float = 1.0

Pooling scales, BayesBlend’s defaults.

adaptive: float

Rate of the exponential prior on lambda (BayesBlend uses 4).

chains: int = 4, 1000, 1000
tune: int = 4, 1000, 1000
draws: int = 4, 1000, 1000
seed: int = 0

Examples

>>> from prospicio.models import HierarchicalStacking
>>> x = [i / 99 - 0.5 for i in range(100)]
>>> a = [-0.5 if v < 0 else -2.0 for v in x]
>>> b = [-2.0 if v < 0 else -0.5 for v in x]
>>> fit = HierarchicalStacking(chains=2, tune=300, draws=300).fit([a, b], [x])
>>> w = fit.weights([[-0.4, 0.4]])
>>> w[0][0] > 0.7 and w[1][0] < 0.3

True

Methods

Name Description
fit() Samples the intercepts and slopes.

fit()

Samples the intercepts and slopes.

Usage

fit(lpd, covariates)
Parameters
lpd: list of list of float

One list per model, one held-out log density per observation.

covariates: list of list of float
One list per covariate, one value per observation.
Returns
StackingFit