superstats.prior#

Prior distributions for simulation parameters.

class superstats.prior.JointPrior(**kwargs)[source]#

Bases: object

Joint prior over multiple model parameters.

Parameters:
**kwargsStochasticTransition, DeterministicTransition, Prior, float, int

Named model parameters.

Use StochasticTransition for stochastic time-varying parameters with hyperparameters, DeterministicTransition for deterministic time-varying parameters, Prior for inferred time-invariant parameters, and scalar values for fixed parameters.

Parameters:

kwargs (StochasticTransition | DeterministicTransition | Prior | float | int)

Notes

Sample outputs are grouped into:

  • local_params: stochastic time-varying parameters (inferred).

  • deterministic_params: deterministic time-varying parameters (no inferred).

  • hyper_params: hyperparameters for transition models (inferred).

  • shared_params: time-invariant parameters (inferred).

  • fixed_params: fixed parameters (no inferred).

plot_joint_prior(num_steps=200, num_trajectories=20, num_draws=1000, **kwargs)[source]#

Plot joint prior diagnostics across local and shared parameters.

Parameters:
num_stepsint, optional, default: 200

Number of time steps for local trajectory sampling.

num_trajectoriesint, optional, default: 20

Number of local trajectories to plot.

num_drawsint, optional, default: 1000

Number of draws used for time-invariant parameter sampling.

**kwargsdict, optional, default: {}

Further optional keyword arguments propagated to the underlying plot_joint_prior plotting function.

Returns:
figplt.Figure - the generated figure
Parameters:
  • num_steps (int)

  • num_trajectories (int)

  • num_draws (int)

plot_time_invariant_prior(num_draws=1000, **kwargs)[source]#

Plot marginal distributions for time-invariant prior parameters.

Parameters:
num_drawsint, optional, default: 1000

Number of draws used to sample hyper_params and shared_params.

**kwargsdict, optional, default: {}

Further optional keyword arguments propagated to the underlying plot_time_invariant_prior plotting function.

Returns:
figplt.Figure - the generated figure
Parameters:

num_draws (int)

plot_time_varying_prior(num_steps=200, num_trajectories=20, **kwargs)[source]#

Plot sampled time-varying prior trajectories.

Parameters:
num_stepsint, optional, default: 200

Number of time steps to sample per trajectory.

num_trajectoriesint, optional, default: 20

Number of trajectories to draw.

**kwargsdict, optional, default: {}

Further optional keyword arguments propagated to the underlying plot_time_varying_prior plotting function.

Returns:
figplt.Figure - the generated figure
Parameters:
  • num_steps (int)

  • num_trajectories (int)

sample(batch_size, num_steps)[source]#

Draw a joint parameter sample.

Parameters:
batch_sizeint

Number of independent samples to draw.

num_stepsint

Number of time steps per trajectory.

Returns:
resultdict - sampled parameter groups local_params,

deterministic_params hyper_params, shared_params, and fixed_params.

Raises:
ValueError

If batch_size or num_steps is not a positive integer.

Parameters:
  • batch_size (int)

  • num_steps (int)

Return type:

Dict[str, Any]

class superstats.prior.Prior(dist, loc=0.0, scale=1.0, low=0.0, high=1.0, a=1.0, b=1.0, alpha=None, scale_factor=1.0, shift=0.0)[source]#

Bases: object

Simple generative prior distribution.

Parameters:
dist{“normal”, “uniform”, “beta”, “halfnormal”, “dirichlet”, “logistic”}

Distribution type.

locfloat, optional, default: 0.0

Mean for normal.

scalefloat, optional, default: 1.0

Standard deviation for normal / halfnormal.

lowfloat, optional, default: 0.0

Lower bound for uniform.

highfloat, optional, default: 1.0

Upper bound for uniform.

afloat, optional, default: 1.0

Alpha (first shape parameter) for beta.

bfloat, optional, default: 1.0

Beta (second shape parameter) for beta.

alphasequence of float or None, optional, default: None

Concentration parameters for dirichlet. Required when dist=”dirichlet” (e.g. [1, 1, 1] for a uniform simplex over 3 categories).

scale_factorfloat, optional, default: 1.0

Multiplicative scaling applied to the drawn samples: scale_factor * samples + shift.

shiftfloat, optional, default: 0.0

Additive offset applied to the drawn samples: scale_factor * samples + shift.

Parameters:
sample(batch_size)[source]#

Draw samples from the prior.

Samples are transformed as scale_factor * samples + shift before being returned.

Parameters:
batch_sizeint

Number of samples to draw.

Returns:
samplesnp.ndarray - shape (batch_size,), or (batch_size, K)

for dirichlet where K is the number of categories

Raises:
ValueError

If dist=”dirichlet” and alpha is None or not a vector-like sequence, or if dist is not one of the supported distributions.

Parameters:

batch_size (int)

Return type:

ndarray

Modules

joint_prior

Joint priors over time-varying and time-invariant parameters.

prior

Elementary prior distributions.