Fit a Bayesian GLM by NUTS
bayes_glm_fit.RdSamples the posterior of a GLM with nuts-rs, the Rust core of nutpie.
The design is R's model.matrix(), as for glm_fit(). Coefficients
have normal priors with mean 0: standard deviation intercept_sd for the
intercept and prior_sd for the rest (on the link scale; standardize
covariates). For the Gaussian, gamma and inverse Gaussian the dispersion
is sampled with a half-normal prior of scale dispersion_scale, unless
dispersion fixes it. Chains run in parallel, start near the
maximum-likelihood fit, and replay exactly from seed. Posterior means
and standard deviations match exact grid integration
(validation/scripts/bayes_glm_grid.py).
Usage
bayes_glm_fit(
formula,
data,
family = "poisson",
link = NULL,
offset = NULL,
weights = NULL,
prior_sd = 2.5,
intercept_sd = 10,
dispersion = NULL,
dispersion_scale = 10,
chains = 4,
tune = 1000,
draws = 1000,
seed = 0,
target_accept = 0.8,
max_depth = 10,
theta = NULL,
power = NULL,
link_power = NULL
)Arguments
- formula
A model formula;
offset(...)terms are honoured.- data
A data frame.
- family
"gaussian","poisson","gamma","inverse_gaussian","binomial"(response a proportion,weightsthe trials),"negative_binomial"(needstheta) or"tweedie"(needspower).- link
"identity","log","logit","probit","cloglog","inverse","inverse_squared"or"power"(needslink_power);NULLfor the family's canonical link.- offset
Optional offset added to any
offset()terms.- weights
Optional prior weights.
- prior_sd, intercept_sd
Prior standard deviations of the slopes and the intercept.
- dispersion
A fixed dispersion, or
NULLfor the family's default.- dispersion_scale
Scale of the half-normal prior on a sampled dispersion.
- chains, tune, draws
Chains, warm-up draws and kept draws per chain.
- seed
Seed, a whole number.
- target_accept
Target acceptance rate for step-size adaptation.
- max_depth
Largest tree depth.
- theta
Negative binomial
theta(variancemu + mu^2 / theta).- power
Tweedie power in
(1, 2).- link_power
Exponent of the power link.
Value
A bayes_glm_model with properties coefficients (posterior
means), summary (a data frame of mean, sd, quantiles, R-hat and ESS
per parameter), draws (a matrix, one row per draw), dispersion_draws
and divergences. Use stats::coef(), stats::predict(),
predict_distribution() and bayes_loo().
Examples
d <- data.frame(claims = rep(c(1, 2, 3, 5), 10), x = rep(c(-1.5, -0.5, 0.5, 1.5), 10))
m <- bayes_glm_fit(claims ~ x, d, family = "poisson", chains = 2, tune = 300, draws = 300)
m@summary
#> parameter mean sd q05 q50 q95 rhat
#> 1 (Intercept) 0.8462963 0.11569769 0.6453650 0.8531218 1.0243032 1.0013364
#> 2 x 0.5078291 0.09628648 0.3417002 0.5136568 0.6601851 0.9983802
#> ess_bulk ess_tail
#> 1 367.4871 327.0323
#> 2 402.9908 457.0260