Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 16 additions & 16 deletions docs/examples/gallery/posterior_sbc.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -28,10 +28,10 @@
"metadata": {},
"outputs": [],
"source": [
"import numpy as np\n",
"import pymc as pm\n",
"from arviz_plots import plot_ecdf_pit, style\n",
"import matplotlib.pyplot as plt\n",
"import numpy as np\n",
"\n",
"import simuk\n",
"\n",
"random_seed = 42\n",
Expand Down Expand Up @@ -96,21 +96,21 @@
"data = np.array([28.0, 8.0, -3.0, 7.0, -1.0, 1.0, 18.0, 12.0])\n",
"sigma = np.array([15.0, 10.0, 16.0, 11.0, 9.0, 11.0, 10.0, 18.0])\n",
"\n",
"coords={\n",
"coords = {\n",
" \"obs\": np.arange(8),\n",
" \"school\": np.arange(8),\n",
" }\n",
"}\n",
"school_idx = np.arange(8)\n",
"\n",
"with pm.Model(coords=coords) as centered_eight:\n",
" school_idx = pm.Data(\"school_idx\", school_idx, dims=\"obs_id\")\n",
" sigma = pm.Data(\"sigma\", sigma, dims=\"obs\")\n",
" data = pm.Data(\"data\", data, dims=\"obs\")\n",
" \n",
" mu = pm.Normal(name='mu', mu=0, sigma=5)\n",
" tau = pm.HalfCauchy('tau', beta=5)\n",
" theta = pm.Normal('theta', mu=mu, sigma=tau, dims=\"school\")\n",
" y_obs = pm.Normal('y', mu=theta[school_idx], sigma=sigma, observed=data, dims=\"obs\")\n"
"\n",
" mu = pm.Normal(name=\"mu\", mu=0, sigma=5)\n",
" tau = pm.HalfCauchy(\"tau\", beta=5)\n",
" theta = pm.Normal(\"theta\", mu=mu, sigma=tau, dims=\"school\")\n",
" y_obs = pm.Normal(\"y\", mu=theta[school_idx], sigma=sigma, observed=data, dims=\"obs\")"
]
},
{
Expand All @@ -129,7 +129,7 @@
"metadata": {},
"outputs": [],
"source": [
"with centered_eight: \n",
"with centered_eight:\n",
" trace = pm.sample(1000, tune=1000, random_seed=random_seed, progressbar=False)"
]
},
Expand All @@ -155,9 +155,7 @@
" with model:\n",
" pm.set_data(\n",
" new_data={\n",
" \"sigma\": np.concatenate(\n",
" [model[\"sigma\"].get_value(), model[\"sigma\"].get_value()]\n",
" ),\n",
" \"sigma\": np.concatenate([model[\"sigma\"].get_value(), model[\"sigma\"].get_value()]),\n",
" \"school_idx\": np.concatenate(\n",
" [model[\"school_idx\"].get_value(), model[\"school_idx\"].get_value()]\n",
" ),\n",
Expand Down Expand Up @@ -245,9 +243,10 @@
}
],
"source": [
"plot_ecdf_pit(sbc.simulations,\n",
" group=\"posterior_sbc\",\n",
" visuals={\"xlabel\": False},\n",
"plot_ecdf_pit(\n",
" sbc.simulations,\n",
" group=\"posterior_sbc\",\n",
" visuals={\"xlabel\": False},\n",
");"
]
},
Expand Down Expand Up @@ -286,6 +285,7 @@
" coords={\"obs\": np.arange(8 + 1)},\n",
" )\n",
"\n",
"\n",
"skewed_sbc = simuk.SBC(\n",
" centered_eight,\n",
" method=\"posterior\",\n",
Expand Down
63 changes: 36 additions & 27 deletions docs/examples/gallery/prior_sbc.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,11 @@
"metadata": {},
"outputs": [],
"source": [
"from arviz_plots import plot_ecdf_pit, style\n",
"import numpy as np\n",
"from arviz_plots import plot_ecdf_pit, style\n",
"\n",
"import simuk\n",
"\n",
"style.use(\"arviz-variat\")"
]
},
Expand Down Expand Up @@ -50,10 +52,10 @@
"sigma = np.array([15.0, 10.0, 16.0, 11.0, 9.0, 11.0, 10.0, 18.0])\n",
"\n",
"with pm.Model() as centered_eight:\n",
" mu = pm.Normal('mu', mu=0, sigma=5)\n",
" tau = pm.HalfCauchy('tau', beta=5)\n",
" theta = pm.Normal('theta', mu=mu, sigma=tau, shape=8)\n",
" y_obs = pm.Normal('y', mu=theta, sigma=sigma, observed=data)"
" mu = pm.Normal(\"mu\", mu=0, sigma=5)\n",
" tau = pm.HalfCauchy(\"tau\", beta=5)\n",
" theta = pm.Normal(\"theta\", mu=mu, sigma=tau, shape=8)\n",
" y_obs = pm.Normal(\"y\", mu=theta, sigma=sigma, observed=data)"
]
},
{
Expand All @@ -70,9 +72,7 @@
"metadata": {},
"outputs": [],
"source": [
"sbc = simuk.SBC(centered_eight,\n",
" num_simulations=100,\n",
" sample_kwargs={'draws': 100, 'tune': 100})\n",
"sbc = simuk.SBC(centered_eight, num_simulations=100, sample_kwargs={\"draws\": 100, \"tune\": 100})\n",
"\n",
"sbc.run_simulations();"
]
Expand Down Expand Up @@ -104,8 +104,9 @@
}
],
"source": [
"plot_ecdf_pit(sbc.simulations,\n",
" visuals={\"xlabel\":False},\n",
"plot_ecdf_pit(\n",
" sbc.simulations,\n",
" visuals={\"xlabel\": False},\n",
");"
]
},
Expand Down Expand Up @@ -147,9 +148,7 @@
"metadata": {},
"outputs": [],
"source": [
"sbc = simuk.SBC(bmb_model,\n",
" num_simulations=100,\n",
" sample_kwargs={'draws': 25, 'tune': 50})\n",
"sbc = simuk.SBC(bmb_model, num_simulations=100, sample_kwargs={\"draws\": 25, \"tune\": 50})\n",
"\n",
"sbc.run_simulations();"
]
Expand Down Expand Up @@ -209,19 +208,20 @@
"source": [
"import numpyro\n",
"import numpyro.distributions as dist\n",
"from jax import random\n",
"from numpyro.infer import NUTS\n",
"\n",
"y = np.array([28.0, 8.0, -3.0, 7.0, -1.0, 1.0, 18.0, 12.0])\n",
"sigma = np.array([15.0, 10.0, 16.0, 11.0, 9.0, 11.0, 10.0, 18.0])\n",
"\n",
"\n",
"def eight_schools_cauchy_prior(J, sigma, y=None):\n",
" mu = numpyro.sample(\"mu\", dist.Normal(0, 5))\n",
" tau = numpyro.sample(\"tau\", dist.HalfCauchy(5))\n",
" with numpyro.plate(\"J\", J):\n",
" theta = numpyro.sample(\"theta\", dist.Normal(mu, tau))\n",
" numpyro.sample(\"y\", dist.Normal(theta, sigma), obs=y)\n",
"\n",
"\n",
"# We use the NUTS sampler\n",
"nuts_kernel = NUTS(eight_schools_cauchy_prior)"
]
Expand All @@ -248,7 +248,8 @@
}
],
"source": [
"sbc = simuk.SBC(nuts_kernel,\n",
"sbc = simuk.SBC(\n",
" nuts_kernel,\n",
" sample_kwargs={\"num_warmup\": 50, \"num_samples\": 75},\n",
" num_simulations=100,\n",
" data_dir={\"J\": 8, \"sigma\": sigma, \"y\": y},\n",
Expand Down Expand Up @@ -283,8 +284,9 @@
}
],
"source": [
"plot_ecdf_pit(sbc.simulations,\n",
" visuals={\"xlabel\":False},\n",
"plot_ecdf_pit(\n",
" sbc.simulations,\n",
" visuals={\"xlabel\": False},\n",
");"
]
},
Expand Down Expand Up @@ -318,10 +320,13 @@
" scale = sigma / np.sqrt(2)\n",
" return {\"y\": rng.laplace(theta, scale)}\n",
"\n",
"sbc = simuk.SBC(centered_eight,\n",
"\n",
"sbc = simuk.SBC(\n",
" centered_eight,\n",
" num_simulations=100,\n",
" simulator=simulator,\n",
" sample_kwargs={'draws': 25, 'tune': 50})\n",
" sample_kwargs={\"draws\": 25, \"tune\": 50},\n",
")\n",
"\n",
"sbc.run_simulations();"
]
Expand Down Expand Up @@ -349,10 +354,10 @@
" scale = sigma / np.sqrt(2)\n",
" return {\"y\": rng.laplace(mu, scale)}\n",
"\n",
"sbc = simuk.SBC(bmb_model,\n",
" num_simulations=100,\n",
" simulator=simulator,\n",
" sample_kwargs={'draws': 25, 'tune': 50})\n",
"\n",
"sbc = simuk.SBC(\n",
" bmb_model, num_simulations=100, simulator=simulator, sample_kwargs={\"draws\": 25, \"tune\": 50}\n",
")\n",
"\n",
"sbc.run_simulations();"
]
Expand Down Expand Up @@ -380,11 +385,13 @@
" scale = sigma / np.sqrt(2)\n",
" return {\"y\": rng.laplace(theta, scale)}\n",
"\n",
"sbc = simuk.SBC(nuts_kernel,\n",
"\n",
"sbc = simuk.SBC(\n",
" nuts_kernel,\n",
" sample_kwargs={\"num_warmup\": 50, \"num_samples\": 75},\n",
" num_simulations=100,\n",
" simulator=simulator,\n",
" data_dir={\"J\": 8, \"sigma\": sigma, \"y\": y}\n",
" data_dir={\"J\": 8, \"sigma\": sigma, \"y\": y},\n",
")\n",
"\n",
"sbc.run_simulations();"
Expand All @@ -407,8 +414,10 @@
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.14.4",
"tags": ["skip-execution"]
"tags": [
"skip-execution"
],
"version": "3.14.4"
}
},
"nbformat": 4,
Expand Down
59 changes: 59 additions & 0 deletions simuk/backend_adapter.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
from abc import ABC, abstractmethod


class BackendAdapter(ABC):
"""Interface every inference-backend adapter must implement for SBC.

Besides the abstract methods below, implementations must expose two
attributes once constructed:

Attributes
----------
var_names : list[str]
Names of the model's free (unobserved) variables. Rank statistics
are computed for these.
observed_vars : list[str]
Names of the model's observed variables. Replicated data is
generated for, and the model re-conditioned on, these.
"""

var_names: list[str]
observed_vars: list[str]

@abstractmethod
def compute_single_rank(self, transform, name, posterior, simulation_idx, ref_params):
pass

@abstractmethod
def get_posterior_predictive_samples(self, num_simulations, seeds, progress_bar):
pass

@abstractmethod
def get_prior_predictive_samples(self, num_samples, seeds):
pass

@abstractmethod
def simulation_params_no_simulator(self, ref_params, predictive):
pass

@abstractmethod
def simulation_params_from_simulator(self, ref_params, predictive):
pass

@abstractmethod
def get_posterior_samples(
self, simulation_parameters, replicated_data, sample_kwargs, seed, method, simulation_idx
):
pass

@abstractmethod
def subsample(self, ref_params, predictive, seed, size):
pass

@abstractmethod
def replicate(self, predictive, idx, simulation_params):
pass

@abstractmethod
def stop_if_cant_run_without_simulator(self):
pass
Loading
Loading