Skip to content

Commit 151f7ab

Browse files
committed
fix(sabc-performance): port sbijax adapter to functional API
1 parent ccc7276 commit 151f7ab

1 file changed

Lines changed: 5 additions & 6 deletions

File tree

experiments/sabc-performance/adapters/sbijax_adapter.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -9,27 +9,26 @@
99
from jax import numpy as jnp
1010
from jax import random as jr
1111

12-
from sbijax import SABC, MultiEps, abs_distance, inference_data_as_dictionary
12+
from sbijax import MultiEps, abs_distance, sabc
1313

1414

1515
def run(name: str, seed: int, budget: dict, out: str) -> None:
1616
"""Run sbijax SABC on task ``name`` and write samples + timing to ``out``."""
1717
prior, simulator, _ = tasks.build_jax_task(name)
1818
observed = jnp.asarray(tasks.load_observed(name))
19-
model = SABC((lambda: prior, simulator), distance_fn=abs_distance)
19+
sampler = sabc(prior, simulator, distance_fn=abs_distance)
2020

2121
def sample_to_numpy(key):
22-
idata, _ = model.sample_posterior(
22+
particles, _ = sampler.sample(
2323
key,
2424
observed,
2525
n_particles=budget["n_particles"],
2626
n_simulation=budget["n_simulation"],
2727
schedule=MultiEps(v=1.0),
2828
)
29-
d = inference_data_as_dictionary(idata.posterior)
3029
cols = [
31-
np.asarray(d[k]).reshape(-1, np.asarray(d[k]).shape[-1])
32-
for k in sorted(d)
30+
np.asarray(particles[k]).reshape(-1, np.asarray(particles[k]).shape[-1])
31+
for k in sorted(particles)
3332
]
3433
return np.concatenate(cols, 1)
3534

0 commit comments

Comments
 (0)