superstats.prior.joint_prior#

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

Classes

JointPrior(**kwargs)

Joint prior over multiple model parameters.

class superstats.prior.joint_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]