bartz.mcmcstep.StepConfig¶
- class bartz.mcmcstep.StepConfig(steps_done, sparse_on_at, resid_reduction_config, count_reduction_config, prec_reduction_config, prec_count_num_trees, sequential_unroll, augment, mesh, leaf_quantization=None)[source]¶
Options for the MCMC step.
- steps_done: Int32[Array, '']¶
The number of MCMC steps completed so far.
- sparse_on_at: Int32[Array, ''] | None¶
After how many steps to turn on variable selection. If
None, variable selection is disabled.
- resid_reduction_config: ReductionConfig¶
How to sum the residuals in each leaf.
- count_reduction_config: ReductionConfig¶
How to count the datapoints in each leaf.
- prec_reduction_config: ReductionConfig¶
How to sum the likelihood precisions in each leaf.
- sequential_unroll: int | bool¶
How much to unroll the sequential accept/reject loop over trees in
step. See theunrollargument ofjax.lax.scan.
- augment: bool¶
Whether to account exactly, via data augmentation, for the decision rules forbidden by the ancestors of each node when updating
Forest.log_s.
- leaf_quantization: Int32[Array, ''] | None = None¶
If set, quantize the leaves to multiples of
eps(resid dtype) * 2 ** leaf_quantizationinState.resid_eff_scaleunits, which makes the running updates ofState.residmostly exact, (almost) stopping their random-walk rounding drift, assuming|resid| < 2 ** (leaf_quantization + 1)holds (in the same units) for most datapoints most of the time. Intended mostly for use with float16 residuals. Sensible settings are 0 and 1, with 0 leaving some drift, 1 practically no drift, and no setting above 1 justifying the reduced accuracy. With enough datapoints or trees the MCMC breaks down because the sampled leaf variation becomes smaller than the quantum, so this setting can not be used liberally.