diff --git a/simuk/numpyro_adapter.py b/simuk/numpyro_adapter.py index 20ca7a5..2cc4eee 100644 --- a/simuk/numpyro_adapter.py +++ b/simuk/numpyro_adapter.py @@ -4,7 +4,8 @@ import jax import numpy as np -from arviz_base import from_numpyro +import xarray as xr +from arviz_base import dict_to_dataset, extract, from_numpyro from numpyro.handlers import seed, trace from numpyro.infer import MCMC, Predictive @@ -24,13 +25,14 @@ def __init__(self, data_dir, numpyro_model, model, simulator, 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"]) + transform(name, posterior[name].isel(sample=i).values) + for i in range(posterior[name].sizes["sample"]) ] ) - return (transformed_posterior < transform(name, ref_params[name][simulation_idx])).sum( - axis=0 - ) + return ( + transformed_posterior + < transform(name, ref_params[name].isel(sample=simulation_idx).values) + ).sum(axis=0) def get_posterior_predictive_samples(self, num_simulations, seeds, progress_bar): raise NotImplementedError("Posterior SBC is not implemented for numpyro") @@ -51,9 +53,15 @@ def get_prior_predictive_samples(self, num_samples, seeds): 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]} + prior_pred = { + key: np.asarray([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} + + prior = dict_to_dataset(prior, sample_dims=["sample"]) + prior_pred = dict_to_dataset(prior_pred, sample_dims=["sample"]) + return prior, prior_pred def _extract_model_info(self, single_seed): @@ -126,14 +134,14 @@ def get_posterior_samples( if k in simulation_parameters.observed_model_vars } mcmc.run(rng_seed, **free_vars_data, **prior_predictive_args) - return from_numpyro(mcmc)["posterior"] + return extract(from_numpyro(mcmc), group="posterior", keep_dataset=True) 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()} + return {k: v.isel(sample=idx).values for k, v in predictive.items()} def stop_if_cant_run_without_simulator(self): if not self.observed_model_vars: @@ -154,4 +162,4 @@ class NumpyroSimulationParams(NamedTuple): observed_vars: list[str] observed_model_vars: list[str] var_names: list[str] - ref_params: dict[str, jax.Array] + ref_params: xr.Dataset diff --git a/simuk/sbc.py b/simuk/sbc.py index 985c8c3..da7e89d 100644 --- a/simuk/sbc.py +++ b/simuk/sbc.py @@ -30,6 +30,7 @@ pass import numpy as np +import xarray as xr from arviz_base import from_dict from tqdm import tqdm @@ -266,7 +267,7 @@ def __init__( self.sample_kwargs = sample_kwargs self.simulations = {name: [] for name in self.adapter.var_names} self._simulations_complete = 0 - self.posteriors = [] + self.posteriors: list[xr.Dataset] = [] self.keep_fits = keep_fits if simulator is not None and not callable(simulator):