diff --git a/docs/examples/gallery/posterior_sbc.ipynb b/docs/examples/gallery/posterior_sbc.ipynb index de77f9a..1caa14b 100644 --- a/docs/examples/gallery/posterior_sbc.ipynb +++ b/docs/examples/gallery/posterior_sbc.ipynb @@ -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", @@ -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\")" ] }, { @@ -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)" ] }, @@ -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", @@ -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", ");" ] }, @@ -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", diff --git a/docs/examples/gallery/prior_sbc.ipynb b/docs/examples/gallery/prior_sbc.ipynb index bb5a924..5ece4ee 100644 --- a/docs/examples/gallery/prior_sbc.ipynb +++ b/docs/examples/gallery/prior_sbc.ipynb @@ -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\")" ] }, @@ -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)" ] }, { @@ -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();" ] @@ -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", ");" ] }, @@ -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();" ] @@ -209,12 +208,12 @@ "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", @@ -222,6 +221,7 @@ " 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)" ] @@ -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", @@ -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", ");" ] }, @@ -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();" ] @@ -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();" ] @@ -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();" @@ -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, diff --git a/simuk/backend_adapter.py b/simuk/backend_adapter.py new file mode 100644 index 0000000..009ffe2 --- /dev/null +++ b/simuk/backend_adapter.py @@ -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 diff --git a/simuk/numpyro_adapter.py b/simuk/numpyro_adapter.py new file mode 100644 index 0000000..20ca7a5 --- /dev/null +++ b/simuk/numpyro_adapter.py @@ -0,0 +1,157 @@ +import inspect +import logging +from typing import NamedTuple + +import jax +import numpy as np +from arviz_base import from_numpyro +from numpyro.handlers import seed, trace +from numpyro.infer import MCMC, Predictive + +from simuk.backend_adapter import BackendAdapter + +log = logging.getLogger(__name__) + + +class NumpyroAdapter(BackendAdapter): + def __init__(self, data_dir, numpyro_model, model, simulator, single_seed): + self.data_dir = data_dir + self.numpyro_model = numpyro_model + self.model = model + self.simulator = simulator + self._extract_model_info(single_seed) + + def compute_single_rank(self, transform, name, posterior, simulation_idx, ref_params): + transformed_posterior = np.array( + [ + transform(name, posterior[name].sel(chain=0).isel(draw=i).values) + for i in range(posterior[name].sizes["draw"]) + ] + ) + return (transformed_posterior < transform(name, ref_params[name][simulation_idx])).sum( + axis=0 + ) + + def get_posterior_predictive_samples(self, num_simulations, seeds, progress_bar): + raise NotImplementedError("Posterior SBC is not implemented for numpyro") + + def get_prior_predictive_samples(self, num_samples, seeds): + """Generate samples to use for the simulations using numpyro.""" + predictive = Predictive(self.model, num_samples=num_samples) + free_vars_data = { + k: v + for k, v in self.data_dir.items() + if k not in self.observed_vars and k in self.model_params + } + samples = predictive(jax.random.PRNGKey(seeds[0]), **free_vars_data) + prior = {k: v for k, v in samples.items() if k not in self.observed_vars} + if self.simulator: + results = [] + for i, vals in enumerate(zip(*prior.values())): + params = dict(zip(prior.keys(), vals)) + params["seed"] = seeds[i] + results.append(self.simulator(**params)) + prior_pred = {key: [result[key] for result in results] for key in results[0]} + else: + prior_pred = {k: v for k, v in samples.items() if k in self.observed_model_vars} + return prior, prior_pred + + def _extract_model_info(self, single_seed): + self.model_params = set(inspect.signature(self.model).parameters.keys()) + with trace() as tr: + with seed(rng_seed=int(single_seed)): + self.numpyro_model.model(**self.data_dir) + self.var_names = [ + name + for name, site in tr.items() + if site["type"] == "sample" and not site.get("is_observed", False) + ] + self.observed_vars = [ + name + for name, site in tr.items() + if site["type"] == "sample" and site.get("is_observed", False) + ] + # Observed model variables are those that are marked as observed + # and are also model function parameters in order to be able to condition on them. + # For instance, this is used to filter out factor variables that are marked as observed + # but cannot be conditioned on. + self.observed_model_vars = [ + name for name in self.observed_vars if name in self.model_params + ] + + def simulation_params_from_simulator(self, ref_params, predictive): + observed_vars = list(predictive.keys()) + observed_model_vars = [name for name in observed_vars if name in self.model_params] + if not observed_model_vars: + raise ValueError("No observed variables to condition on") + + return NumpyroSimulationParams( + observed_vars=observed_vars, + observed_model_vars=observed_model_vars, + var_names=list( + filter( + lambda var_name: var_name not in observed_vars, + list(ref_params.keys()), + ) + ), + ref_params=ref_params, + ) + + def simulation_params_no_simulator(self, ref_params, predictive): + return NumpyroSimulationParams( + observed_vars=self.observed_vars, + observed_model_vars=self.observed_model_vars, + var_names=self.var_names, + ref_params=ref_params, + ) + + def get_posterior_samples( + self, simulation_parameters, replicated_data, sample_kwargs, seed, method, simulation_idx + ): + """Generate posterior samples using numpyro conditioned to a prior predictive sample.""" + if method == "posterior": + raise NotImplementedError("Posterior SBC not implemented for numpyro") + + mcmc = MCMC(self.numpyro_model, **sample_kwargs) + rng_seed = jax.random.PRNGKey(seed) + + free_vars_data = { + k: v + for k, v in self.data_dir.items() + if k not in simulation_parameters.observed_model_vars and k in self.model_params + } + prior_predictive_args = { + k: v + for k, v in replicated_data.items() + if k in simulation_parameters.observed_model_vars + } + mcmc.run(rng_seed, **free_vars_data, **prior_predictive_args) + return from_numpyro(mcmc)["posterior"] + + def subsample(self, ref_params, predictive, seed, size): + log.info("Slicing isn't implemented for numpyro, skipping it.") + return ref_params, predictive + + def replicate(self, predictive, idx, simulation_params): + return {k: v[idx] for k, v in predictive.items()} + + def stop_if_cant_run_without_simulator(self): + if not self.observed_model_vars: + raise ValueError( + "There are no observed variables we can condition on, and NumPyro " + "will not generate prior predictive samples. Either change the model " + "or specify a simulator with the `simulator` argument." + ) + missing = [name for name in self.observed_model_vars if name not in self.data_dir] + if missing: + raise ValueError( + "The following model parameters are missing from data_dir: " + + ", ".join(sorted(missing)) + ) + + +class NumpyroSimulationParams(NamedTuple): + observed_vars: list[str] + observed_model_vars: list[str] + var_names: list[str] + ref_params: dict[str, jax.Array] diff --git a/simuk/pymc_adapter.py b/simuk/pymc_adapter.py new file mode 100644 index 0000000..91dd553 --- /dev/null +++ b/simuk/pymc_adapter.py @@ -0,0 +1,261 @@ +import traceback +from collections.abc import Mapping +from typing import NamedTuple + +import numpy as np +import pymc as pm +import xarray as xr +from arviz_base import dict_to_dataset, extract + +from simuk.backend_adapter import BackendAdapter + + +class PymcAdapter(BackendAdapter): + def __init__(self, model, simulator, trace, augment_observed, update_data): + self.model = model + self.simulator = simulator + self.trace = trace + self.augment_observed = augment_observed + self.update_data = update_data + self._extract_model_info() + + def get_posterior_predictive_samples(self, num_simulations, seeds, progress_bar): + with self.model: + num_draws = self.trace["posterior"].sizes["draw"] + draw_indices = np.linspace(0, num_draws - 1, num_simulations, dtype=int) + thinned_idata = self.trace.isel(draw=draw_indices) + posterior = extract(thinned_idata, group="posterior", keep_dataset=True) + + if self.simulator is None: + pm.sample_posterior_predictive( + thinned_idata, + extend_inferencedata=True, + random_seed=seeds[0], + progressbar=progress_bar, + ) + posterior_pred = extract( + thinned_idata, group="posterior_predictive", keep_dataset=True + ) + return posterior, posterior_pred + else: + posterior_pred = self._get_simulator_data(posterior, seeds) + + return posterior, posterior_pred + + def compute_single_rank(self, transform, name, posterior, simulation_idx, ref_params): + transformed_posterior = np.array( + [ + transform(name, posterior[name].isel(sample=i).values) + for i in range(posterior[name].sizes["sample"]) + ] + ) + return ( + transformed_posterior + < transform(name, ref_params[name].isel(sample=simulation_idx).values) + ).sum(axis=0) + + def get_prior_predictive_samples(self, num_samples, seeds): + """Generate samples to use for the simulations.""" + with self.model: + idata = pm.sample_prior_predictive(draws=num_samples, random_seed=seeds[0]) + prior = extract(idata, group="prior", keep_dataset=True) + + if self.simulator is None: + prior_pred = extract(idata, group="prior_predictive", keep_dataset=True) + return prior, prior_pred + + prior_pred = self._get_simulator_data(prior, seeds) + + return prior, prior_pred + + def _get_simulator_data(self, free_rv_samples, seeds): + """Run the user-defined simulator to obtain predictive samples. + + These samples can be generated from either prior or posterior samples. + """ + # Deal with custom simulator + pred = [] + for i in range(free_rv_samples.sizes["sample"]): + params = { + var: free_rv_samples[var].isel(sample=i).values for var in free_rv_samples.data_vars + } + params["seed"] = seeds[i] + try: + res = self.simulator(**params) + except Exception as e: + raise ValueError( + f"Error generating prior predictive sample with parameters {params}: {e}." + ) from e + + if not isinstance(res, Mapping): + raise TypeError(f"Simulator must return a dictionary, got {type(res)}") + + pred.append(res) + + pred = dict_to_dataset( + {key: np.stack([pp[key] for pp in pred]) for key in pred[0]}, + sample_dims=["sample"], + coords={**free_rv_samples.coords}, + ) + + return pred + + def _extract_model_info(self): + """Extract observed and free variables from the model. + + Also records the baseline state for Posterior SBC. + """ + observed_var_nodes = [obs_rv for obs_rv in self.model.observed_RVs] + self.observed_vars = [obs.name for obs in observed_var_nodes] + self.var_names = [v.name for v in self.model.free_RVs] + # Stores what observed values are given by pm.Data + self.observed_rvs_to_pm_data = { + var.name: ( + self.model.rvs_to_values[var].name + if hasattr(self.model.rvs_to_values[var], "get_value") + else None + ) + for var in observed_var_nodes + } + self.model_baseline_state = self._get_baseline_state(self.model) + + def _get_baseline_state(self, model): + """Extract the current mutable data and coordinates from a PyMC model.""" + baseline_data = {} + + # Extract Mutable Data + for var in model.data_vars: + if hasattr(var, "get_value"): + baseline_data[var.name] = var.get_value(borrow=False) + + # Extract Coordinates + # Convert the internal PyMC coordinate object to a standard dictionary + baseline_coords = dict(model.coords) + + return {"data": baseline_data, "coords": baseline_coords} + + def simulation_params_no_simulator(self, ref_params, predictive): + return PymcSimulationParams( + observed_vars=self.observed_vars, var_names=self.var_names, ref_params=ref_params + ) + + def simulation_params_from_simulator(self, ref_params, predictive): + observed_vars = list(predictive.data_vars) + return PymcSimulationParams( + observed_vars=observed_vars, + var_names=list( + filter( + lambda var_name: var_name not in observed_vars, + list(ref_params.data_vars), + ) + ), + ref_params=ref_params, + ) + + def subsample(self, ref_params, predictive, seed, size): + rng = np.random.default_rng(seed) + sample_indices = rng.choice(ref_params.sizes["sample"], size=size, replace=False) + ref_params = ref_params.isel(sample=sample_indices) + predictive = predictive.isel(sample=sample_indices) + return ref_params, predictive + + def replicate(self, predictive, idx, simulation_params): + return { + var_name: predictive[var_name].isel(sample=idx).values + for var_name in simulation_params.observed_vars + } + + def get_posterior_samples( + self, simulation_parameters, replicated_data, sample_kwargs, seed, method, simulation_idx + ): + """Fit the model and return posterior draws for one SBC iteration. + + For **Prior SBC** the model is conditioned on the replicated data + alone. For **Posterior SBC** the original observed data and the + replicated data are combined (via ``augment_observed`` or the default + simple concatenation) and the model is conditioned on the augmented + dataset. + + Parameters + ---------- + replicated_data : dict[str, np.ndarray] + Simulated observations for the current iteration, keyed by + observed-variable name. + + Returns + ------- + xarray.Dataset + Posterior draws from the (augmented) model. + """ + if method == "posterior": + observed_data = self.trace["observed_data"] + + if self.augment_observed is not None: + augmented_data = self.augment_observed( + self.model, observed_data, replicated_data, simulation_idx + ) + else: + # Default: concatenate original and replicated observations + augmented_data = { + var_name: np.concatenate( + [observed_data[var_name].values, replicated_data[var_name]] + ) + for var_name in simulation_parameters.observed_vars + } + + if self.update_data is not None: + with self.model: + self.update_data(self.model, augmented_data, simulation_idx) + + vars_to_observations = augmented_data + else: + # Prior SBC simply uses the generated prior predictive replicated data + vars_to_observations = replicated_data + + # Set observed data that are pm.Data objects if the user hasn't modified them yet. + # We enforce an np.array_equal check against the baseline to prevent PyMC size mismatch + # ValueErrors when the user's `update_data` hook or `pm.observe` already updated it. + with self.model: + for rv, data_node in self.observed_rvs_to_pm_data.items(): + if data_node is not None and np.array_equal( + self.model.named_vars[data_node].get_value(), + self.model_baseline_state["data"][data_node], + ): + pm.set_data(new_data={data_node: vars_to_observations[rv]}) + + try: + new_model = pm.observe(self.model, vars_to_observations=vars_to_observations) + with new_model: + check = pm.sample(**sample_kwargs, random_seed=seed) + + posterior = extract(check, group="posterior", keep_dataset=True) + except Exception: + traceback.print_exc() + raise + finally: + # Always ensure the model is reset to its un-augmented baseline state + # so the next simulation iteration isn't corrupted by the previous loop's augmented data + self._reset_model_state(self.model, self.model_baseline_state) + + return posterior + + def _reset_model_state(self, model, model_state): + """Reset the state of PyMC model.""" + with model: + pm.set_data(model_state["data"], coords=model_state["coords"]) + + def stop_if_cant_run_without_simulator(self): + if not self.observed_vars: + # Ideally, we could raise an error early for `numpyro` also, + # but `factor` also produces 'observed_vars' + raise ValueError( + "There are no observed variables, and PyMC will not generate predictive " + "samples for both Prior and Posterior SBC. Either change the model or " + "specify a simulator with the `simulator` argument." + ) + + +class PymcSimulationParams(NamedTuple): + observed_vars: list[str] + var_names: list[str] + ref_params: xr.Dataset diff --git a/simuk/sbc.py b/simuk/sbc.py index e678d38..985c8c3 100644 --- a/simuk/sbc.py +++ b/simuk/sbc.py @@ -14,29 +14,27 @@ """ import logging -import traceback from copy import copy from importlib.metadata import version +# Both backends are optional: these imports only provide the names used by the +# engine-detection checks in SBC.__init__, which short-circuit before touching +# a name whose backend is not installed. try: import pymc as pm except ImportError: pass try: - import jax - from numpyro.handlers import seed, trace - from numpyro.infer import MCMC, Predictive from numpyro.infer.mcmc import MCMCKernel except ImportError: pass -import inspect -from collections.abc import Mapping - import numpy as np -from arviz_base import dict_to_dataset, extract, from_dict, from_numpyro +from arviz_base import from_dict from tqdm import tqdm +_log = logging.getLogger(__name__) + class quiet_logging: """Turn off logging for PyMC, Bambi and PyTensor.""" @@ -215,22 +213,37 @@ def __init__( keep_fits=True, progress_bar=True, ): + self.num_simulations = num_simulations + self.seed = seed + self._seeds = self._get_seeds() + if hasattr(model, "basic_RVs") and isinstance(model, pm.Model): + from simuk.pymc_adapter import PymcAdapter # noqa: PLC0415 + self.engine = "pymc" self.model = model + self.adapter = PymcAdapter(self.model, simulator, trace, augment_observed, update_data) elif hasattr(model, "formula"): + from simuk.pymc_adapter import PymcAdapter # noqa: PLC0415 + self.engine = "bambi" model.build() self.bambi_model = model self.model = model.backend.model self.formula = model.formula self.new_data = copy(model.data) + self.adapter = PymcAdapter(self.model, simulator, trace, augment_observed, update_data) elif isinstance(model, MCMCKernel): + # runtime import so an environment with only Pymc can run SBC over Pymc models. + from simuk.numpyro_adapter import NumpyroAdapter # noqa: PLC0415 + self.engine = "numpyro" self.numpyro_model = model self.model = self.numpyro_model.model - self.run_simulations = self._run_simulations_numpyro self.data_dir = data_dir if data_dir is not None else {} + self.adapter = NumpyroAdapter( + self.data_dir, self.numpyro_model, self.model, simulator, self._seeds[0] + ) else: raise ValueError( "model should be one of pymc.Model, bambi.Model, or numpyro.infer.mcmc.MCMCKernel" @@ -251,47 +264,21 @@ def __init__( sample_kwargs.setdefault("progressbar", False) sample_kwargs.setdefault("compute_convergence_checks", False) self.sample_kwargs = sample_kwargs - - self.num_simulations = num_simulations - self.seed = seed - self._seeds = self._get_seeds() - - self._extract_model_info() - self.simulations = {name: [] for name in self.var_names} + self.simulations = {name: [] for name in self.adapter.var_names} self._simulations_complete = 0 self.posteriors = [] self.keep_fits = keep_fits - self.ref_params = None if simulator is not None and not callable(simulator): raise ValueError("simulator should be a function or None") - if simulator is not None and self.observed_vars: + if simulator is not None and self.adapter.observed_vars: logging.warning( "Provided model contains both observed variables and a simulator. " "Ignoring observed variables and using the simulator instead." ) - if simulator is None and not self.observed_vars and self.engine == "pymc": - # Ideally, we could raise an error early for `numpyro` also, - # but `factor` also produces 'observed_vars' - raise ValueError( - "There are no observed variables, and PyMC will not generate predictive " - "samples for both Prior and Posterior SBC. Either change the model or " - "specify a simulator with the `simulator` argument." - ) + if simulator is None: + self.adapter.stop_if_cant_run_without_simulator() - if simulator is None and self.engine == "numpyro": - if not self.observed_model_vars: - raise ValueError( - "There are no observed variables we can condition on, and NumPyro " - "will not generate prior predictive samples. Either change the model " - "or specify a simulator with the `simulator` argument." - ) - missing = [name for name in self.observed_model_vars if name not in self.data_dir] - if missing: - raise ValueError( - "The following model parameters are missing from data_dir: " - + ", ".join(sorted(missing)) - ) self.simulator = simulator self._transform = lambda param_name, param_value: param_value @@ -343,256 +330,11 @@ def __init__( if trace is not None: logging.warning("`trace` is only used for Posterior SBC. Ignoring...") - def _extract_model_info(self): - """Extract observed and free variables from the model. - - Also records the baseline state for Posterior SBC. - """ - if self.engine == "numpyro": - self.model_params = set(inspect.signature(self.model).parameters.keys()) - with trace() as tr: - with seed(rng_seed=int(self._seeds[0])): - self.numpyro_model.model(**self.data_dir) - self.var_names = [ - name - for name, site in tr.items() - if site["type"] == "sample" and not site.get("is_observed", False) - ] - self.observed_vars = [ - name - for name, site in tr.items() - if site["type"] == "sample" and site.get("is_observed", False) - ] - # Observed model variables are those that are marked as observed - # and are also model function parameters in order to be able to condition on them. - # For instance, this is used to filter out factor variables that are marked as observed - # but cannot be conditioned on. - self.observed_model_vars = [ - name for name in self.observed_vars if name in self.model_params - ] - - else: - observed_var_nodes = [obs_rv for obs_rv in self.model.observed_RVs] - self.observed_vars = [obs.name for obs in observed_var_nodes] - self.var_names = [v.name for v in self.model.free_RVs] - # Stores what observed values are given by pm.Data - self.observed_rvs_to_pm_data = { - var.name: ( - self.model.rvs_to_values[var].name - if hasattr(self.model.rvs_to_values[var], "get_value") - else None - ) - for var in observed_var_nodes - } - self.model_baseline_state = self._get_baseline_state(self.model) - - def _get_baseline_state(self, model): - """Extract the current mutable data and coordinates from a PyMC model.""" - baseline_data = {} - - # Extract Mutable Data - for var in model.data_vars: - if hasattr(var, "get_value"): - baseline_data[var.name] = var.get_value(borrow=False) - - # Extract Coordinates - # Convert the internal PyMC coordinate object to a standard dictionary - baseline_coords = dict(model.coords) - - return {"data": baseline_data, "coords": baseline_coords} - - def _reset_model_state(self, model, model_state): - """Reset the state of PyMC model.""" - with model: - pm.set_data(model_state["data"], coords=model_state["coords"]) - def _get_seeds(self): """Set the random seed, and generate seeds for all the simulations.""" rng = np.random.default_rng(self.seed) return rng.integers(0, 2**30, size=self.num_simulations) - def _get_simulator_data(self, free_rv_samples): - """Run the user-defined simulator to obtain predictive samples. - - These samples can be generated from either prior or posterior samples. - """ - # Deal with custom simulator - pred = [] - for i in range(free_rv_samples.sizes["sample"]): - params = { - var: free_rv_samples[var].isel(sample=i).values for var in free_rv_samples.data_vars - } - params["seed"] = self._seeds[i] - try: - res = self.simulator(**params) - except Exception as e: - raise ValueError( - f"Error generating prior predictive sample with parameters {params}: {e}." - ) - - if not isinstance(res, Mapping): - raise TypeError(f"Simulator must return a dictionary, got {type(res)}") - - pred.append(res) - - pred = dict_to_dataset( - {key: np.stack([pp[key] for pp in pred]) for key in pred[0]}, - sample_dims=["sample"], - coords={**free_rv_samples.coords}, - ) - - return pred - - def _get_prior_predictive_samples(self): - """Generate samples to use for the simulations.""" - with self.model: - idata = pm.sample_prior_predictive( - draws=self.num_simulations, random_seed=self._seeds[0] - ) - prior = extract(idata, group="prior", keep_dataset=True) - - if self.simulator is None: - prior_pred = extract(idata, group="prior_predictive", keep_dataset=True) - return prior, prior_pred - - prior_pred = self._get_simulator_data(prior) - - return prior, prior_pred - - def _get_prior_predictive_samples_numpyro(self): - """Generate samples to use for the simulations using numpyro.""" - predictive = Predictive(self.model, num_samples=self.num_simulations) - free_vars_data = { - k: v - for k, v in self.data_dir.items() - if k not in self.observed_vars and k in self.model_params - } - samples = predictive(jax.random.PRNGKey(self._seeds[0]), **free_vars_data) - prior = {k: v for k, v in samples.items() if k not in self.observed_vars} - if self.simulator: - results = [] - for i, vals in enumerate(zip(*prior.values())): - params = dict(zip(prior.keys(), vals)) - params["seed"] = self._seeds[i] - results.append(self.simulator(**params)) - prior_pred = {key: [result[key] for result in results] for key in results[0]} - else: - prior_pred = {k: v for k, v in samples.items() if k in self.observed_model_vars} - return prior, prior_pred - - def _get_posterior_samples(self, replicated_data): - """Fit the model and return posterior draws for one SBC iteration. - - For **Prior SBC** the model is conditioned on the replicated data - alone. For **Posterior SBC** the original observed data and the - replicated data are combined (via ``augment_observed`` or the default - simple concatenation) and the model is conditioned on the augmented - dataset. - - Parameters - ---------- - replicated_data : dict[str, np.ndarray] - Simulated observations for the current iteration, keyed by - observed-variable name. - - Returns - ------- - xarray.Dataset - Posterior draws from the (augmented) model. - """ - if self.method == "posterior": - observed_data = self.trace["observed_data"] - - if self.augment_observed is not None: - augmented_data = self.augment_observed( - self.model, observed_data, replicated_data, self._simulations_complete - ) - else: - # Default: concatenate original and replicated observations - augmented_data = { - var_name: np.concatenate( - [observed_data[var_name].values, replicated_data[var_name]] - ) - for var_name in self.observed_vars - } - - if self.update_data is not None: - with self.model: - self.update_data(self.model, augmented_data, self._simulations_complete) - - vars_to_observations = augmented_data - else: - # Prior SBC simply uses the generated prior predictive replicated data - vars_to_observations = replicated_data - - # Set observed data that are pm.Data objects if the user hasn't modified them yet. - # We enforce an np.array_equal check against the baseline to prevent PyMC size mismatch - # ValueErrors when the user's `update_data` hook or `pm.observe` already updated it. - with self.model: - for rv, data_node in self.observed_rvs_to_pm_data.items(): - if data_node is not None and np.array_equal( - self.model.named_vars[data_node].get_value(), - self.model_baseline_state["data"][data_node], - ): - pm.set_data(new_data={data_node: vars_to_observations[rv]}) - - try: - new_model = pm.observe(self.model, vars_to_observations=vars_to_observations) - with new_model: - check = pm.sample( - **self.sample_kwargs, random_seed=self._seeds[self._simulations_complete] - ) - - posterior = extract(check, group="posterior", keep_dataset=True) - except Exception: - traceback.print_exc() - raise - finally: - # Always ensure the model is reset to its un-augmented baseline state - # so the next simulation iteration isn't corrupted by the previous loop's augmented data - self._reset_model_state(self.model, self.model_baseline_state) - - return posterior - - def _get_posterior_samples_numpyro(self, prior_predictive_draw): - """Generate posterior samples using numpyro conditioned to a prior predictive sample.""" - mcmc = MCMC(self.numpyro_model, **self.sample_kwargs) - rng_seed = jax.random.PRNGKey(self._seeds[self._simulations_complete]) - - free_vars_data = { - k: v - for k, v in self.data_dir.items() - if k not in self.observed_model_vars and k in self.model_params - } - prior_predictive_args = { - k: v for k, v in prior_predictive_draw.items() if k in self.observed_model_vars - } - mcmc.run(rng_seed, **free_vars_data, **prior_predictive_args) - return from_numpyro(mcmc)["posterior"] - - def _get_posterior_predictive_samples(self): - with self.model: - num_draws = self.trace["posterior"].sizes["draw"] - draw_indices = np.linspace(0, num_draws - 1, self.num_simulations, dtype=int) - thinned_idata = self.trace.isel(draw=draw_indices) - posterior = extract(thinned_idata, group="posterior", keep_dataset=True) - - if self.simulator is None: - pm.sample_posterior_predictive( - thinned_idata, - extend_inferencedata=True, - random_seed=self._seeds[0], - progressbar=self.progress_bar, - ) - posterior_pred = extract( - thinned_idata, group="posterior_predictive", keep_dataset=True - ) - return posterior, posterior_pred - else: - posterior_pred = self._get_simulator_data(posterior) - - return posterior, posterior_pred - def _convert_to_datatree(self): """Pack the rank-statistic arrays into an xarray DataTree. @@ -648,45 +390,24 @@ def compute_rank_statistics(self, transform=None): elif not callable(transform): raise ValueError("`transform` should be a function or None") - self.simulations = {name: [] for name in self.var_names} + self.simulations = {name: [] for name in self.kept_simulation_params.var_names} for idx, posterior in enumerate(self.posteriors): - self._compute_single_rank(idx, posterior, transform) + self._compute_single_rank(idx, posterior, transform, self.kept_simulation_params) self.simulations = {k: np.stack(v)[None, :] for k, v in self.simulations.items()} self._convert_to_datatree() return self.simulations - def _compute_single_rank(self, simulation_idx, posterior, transform): - for name in self.var_names: - if self.engine == "numpyro": - transformed_posterior = np.array( - [ - transform(name, posterior[name].sel(chain=0).isel(draw=i).values) - for i in range(posterior[name].sizes["draw"]) - ] - ) - self.simulations[name].append( - ( - transformed_posterior - < transform(name, self.ref_params[name][simulation_idx]) - ).sum(axis=0) - ) - elif self.engine in ["bambi", "pymc"]: - transformed_posterior = np.array( - [ - transform(name, posterior[name].isel(sample=i).values) - for i in range(posterior[name].sizes["sample"]) - ] - ) - self.simulations[name].append( - ( - transformed_posterior - < transform(name, self.ref_params[name].isel(sample=simulation_idx).values) - ).sum(axis=0) + def _compute_single_rank(self, simulation_idx, posterior, transform, simulation_params): + for name in simulation_params.var_names: + self.simulations[name].append( + self.adapter.compute_single_rank( + transform, name, posterior, simulation_idx, simulation_params.ref_params ) + ) - @quiet_logging("pymc", "pytensor.gof.compilelock", "bambi") + @quiet_logging("pymc", "pytensor.gof.compilelock", "bambi", "numpyro") def run_simulations(self): """Run all SBC iterations (Prior or Posterior SBC). @@ -706,6 +427,10 @@ def run_simulations(self): you can keyboard-interrupt part way through, inspect the partial results, and then call ``run_simulations()`` again to continue. If a seed was passed at init, reproducibility is preserved. + + If an error occurs during a simulation, it is logged (with its + traceback) rather than raised, and the run finalizes with the + rank statistics of the iterations completed so far. """ progress = tqdm( initial=self._simulations_complete, @@ -716,99 +441,55 @@ def run_simulations(self): if self.method == "prior": # In Prior SBC, the reference parameter draws are from the prior, # the predictive samples are from the prior predictive - ref_params, predictive = self._get_prior_predictive_samples() + ref_params, predictive = self.adapter.get_prior_predictive_samples( + self.num_simulations, self._seeds + ) else: # In Posterior SBC, the reference parameter draws are from the original posterior, # the predictive samples are from the original posterior predictive - ref_params, predictive = self._get_posterior_predictive_samples() + ref_params, predictive = self.adapter.get_posterior_predictive_samples( + self.num_simulations, self._seeds, self.progress_bar + ) - rng = np.random.default_rng(self.seed) - sample_indices = rng.choice( - ref_params.sizes["sample"], size=self.num_simulations, replace=False + ref_params, predictive = self.adapter.subsample( + ref_params, predictive, self.seed, self.num_simulations ) - self.ref_params = ref_params.isel(sample=sample_indices) - predictive = predictive.isel(sample=sample_indices) - # if simulator is used, ignore observed_vars if self.simulator is not None: - self.observed_vars = list(predictive.data_vars) - self.var_names = list( - filter( - lambda var_name: var_name not in self.observed_vars, - list(ref_params.data_vars), - ) + # if simulator is used, ignore observed_vars + simulation_params = self.adapter.simulation_params_from_simulator( + ref_params, predictive ) - self.simulations = {var_name: [] for var_name in self.var_names} + self.simulations = {var_name: [] for var_name in simulation_params.var_names} + else: + simulation_params = self.adapter.simulation_params_no_simulator(ref_params, predictive) + + self._simulation_loop(progress, simulation_params, predictive) + def _simulation_loop(self, progress, simulation_params, predictive): try: while self._simulations_complete < self.num_simulations: idx = self._simulations_complete - replicated_data = { - var_name: predictive[var_name].isel(sample=idx).values - for var_name in self.observed_vars - } - - posterior = self._get_posterior_samples(replicated_data) - if self.keep_fits: - self.posteriors.append(posterior) - else: - self._compute_single_rank(idx, posterior, self._transform) - - self._simulations_complete += 1 - progress.update() - except Exception: - logging.error("Stopping simulation. An error occurred during simulations:") - traceback.print_exc() - finally: - if self._simulations_complete: - if self.keep_fits: - self.compute_rank_statistics() - else: - self.simulations = { - k: np.stack(v)[None, :] for k, v in self.simulations.items() - } - self._convert_to_datatree() - - progress.close() - - @quiet_logging("numpyro") - def _run_simulations_numpyro(self): - """Run all the simulations for Numpyro Model.""" - prior, prior_pred = self._get_prior_predictive_samples_numpyro() - self.ref_params = prior - progress = tqdm( - initial=self._simulations_complete, - total=self.num_simulations, - ) - # if simulator is used, ignore observed_vars - if self.simulator is not None: - self.observed_vars = list(prior_pred.keys()) - self.observed_model_vars = [ - name for name in self.observed_vars if name in self.model_params - ] - if not self.observed_model_vars: - raise ValueError("No observed variables to condition on") - - self.var_names = list( - filter( - lambda var_name: var_name not in self.observed_vars, - list(prior.keys()), + replicated_data = self.adapter.replicate(predictive, idx, simulation_params) + posterior = self.adapter.get_posterior_samples( + simulation_params, + replicated_data, + self.sample_kwargs, + self._seeds[self._simulations_complete], + self.method, + self._simulations_complete, ) - ) - self.simulations = {var_name: [] for var_name in self.var_names} - try: - while self._simulations_complete < self.num_simulations: - idx = self._simulations_complete - prior_predictive_draw = {k: v[idx] for k, v in prior_pred.items()} - posterior = self._get_posterior_samples_numpyro(prior_predictive_draw) if self.keep_fits: self.posteriors.append(posterior) + self.kept_simulation_params = simulation_params else: - self._compute_single_rank(idx, posterior, self._transform) + self._compute_single_rank(idx, posterior, self._transform, simulation_params) self._simulations_complete += 1 progress.update() + except Exception: + _log.exception("Stopping simulation. An error occurred during simulations:") finally: if self._simulations_complete: if self.keep_fits: diff --git a/simuk/tests/test_prior_sbc.py b/simuk/tests/test_prior_sbc.py index 092cd41..4c3a545 100644 --- a/simuk/tests/test_prior_sbc.py +++ b/simuk/tests/test_prior_sbc.py @@ -254,6 +254,28 @@ def test_compute_rank_statistics_transform_not_callable(): sbc.compute_rank_statistics(transform=123) +def test_compute_rank_statistics_recompute_with_new_transform(): + sbc = simuk.SBC( + centered_eight, + num_simulations=2, + sample_kwargs={"draws": 5, "tune": 5}, + ) + sbc.run_simulations() + # With the default identity transform, theta keeps its vector shape + assert sbc.simulations["prior_sbc"]["theta"].shape[-1] == 8 + + num_posteriors = len(sbc.posteriors) + recomputed = sbc.compute_rank_statistics( + transform=lambda param_name, param_value: np.mean(param_value) + ) + assert "prior_sbc" in recomputed + # The mean transform reduces the vector parameter to a scalar test quantity + assert recomputed["prior_sbc"]["theta"].shape == (1, 2) + # Recomputation reuses the stored fits instead of rerunning simulations + assert len(sbc.posteriors) == num_posteriors + assert sbc._simulations_complete == 2 + + def test_sbc_run_simulations_keep_fits_false(): sbc = simuk.SBC( centered_eight,