# Use of JAX-based generative models - explicit RNG management

**URL:** https://discuss.bayesflow.org/t/use-of-jax-based-generative-models-explicit-rng-management/241
**Category:** General
**Created:** [April 28, 2026, 7:03am UTC](https://discuss.bayesflow.org/t/use-of-jax-based-generative-models-explicit-rng-management/241 "2026-04-28T07:03:39Z")
**Posts on this page:** 3
**Page:** 1

<div class="post-metadata">

### Author: ![floralynth](https://avatars.discourse-cdn.com/v4/letter/f/3bc359/32.png) [@floralynth](https://discuss.bayesflow.org/u/floralynth)
#### Post date: [April 28, 2026, 7:03am UTC](https://discuss.bayesflow.org/t/use-of-jax-based-generative-models-explicit-rng-management/241/1 "2026-04-28T07:03:39Z")

</div>

Hi all,

First time here and aiming to use Bayesflow to speed up some existing MCMC based workflows for cognitive models.

I already have a series of models written using JAX-numpy and numpyro which can be used for likelihood and simulation. Ideally, I would simply adapt/wrap these models to use them with Bayesflow, but they would require explicit RNG management.

So I would ideally be able to define my functions as

key = jax.random.PRNGKey(3)  
def model\_for\_simulation(rng\_key = key, simulate = True):  
key, key2 = jax.random.split(key)  
sims = dict(a = dist.Normal(0, 1).sample(key2))

return sims, key

Or something like that. It seems however, that despite using JAX as a backend, there isn’t really a way to interface with RNG management in the ways required to use it as a modelling language?

Any ideas? perhaps I could create a wrapper that manages the key? My only concern is that I am not sure if that would affect Bayesflows assumptions about independence across batch dims etc.

Cheers 🙂

---

<div class="post-metadata">

### Author: ![hanol](https://yyz1.discourse-cdn.com/flex007/user_avatar/discuss.bayesflow.org/hanol/32/165_2.png) [@hanol](https://discuss.bayesflow.org/u/hanol)
#### Post date: [April 28, 2026, 11:42am UTC](https://discuss.bayesflow.org/t/use-of-jax-based-generative-models-explicit-rng-management/241/2 "2026-04-28T11:42:22Z")

</div>

Hi and welcome!

Conventions regarding RNG across packages are unfortunately not very standardized, but we are trying to improve this.

BayesFlow is very permissive in terms of simulator randomness:

- most straightforwardly, you can always generate simulations with `jax.random.PRNGKey` as you mention in your message, then train with `fit_offline(data=sims, ...)`.
- in case you want to use `fit_online`, you need to pass a callable simulator object taking a batch\_size/shape, returning a dict. You are free to use any specific RNG inside of that callable.

Are these applicable? If not, please could you clarify what your workflow requires in terms of RNG management?

Further:

- The next BayesFlow release will expose a seed argument to all network `sample` methods which can be either a fixed seed or a backend-agnostic keras [seed generator](https://keras.io/api/random/seed_generator/). It is already on the [dev](https://github.com/bayesflow-org/bayesflow/tree/dev) branch.
- In case you’re interested, there is currently an ONNX working group aiming to standardize RNG for probabilistic programming: [working-groups/probabilistic-programming at main · onnx/working-groups · GitHub](https://github.com/onnx/working-groups/tree/main/probabilistic-programming).

Interested how your efforts speeding up MCMC for your models are going - feel free to reach out anytime!

---

<div class="post-metadata">

### Author: ![floralynth](https://avatars.discourse-cdn.com/v4/letter/f/3bc359/32.png) [@floralynth](https://discuss.bayesflow.org/u/floralynth)
#### Post date: [May 5, 2026, 6:25am UTC](https://discuss.bayesflow.org/t/use-of-jax-based-generative-models-explicit-rng-management/241/3 "2026-05-05T06:25:11Z")

</div>

Hey, letting you know that I was able to work around this issue by only jax.jit compiling functions AFTER splitting keys in pure python, then passing a vector of keys to a jit compiled function. Probably a bit slower than the absolute best case scenario but with minimal rewriting.

I also suspect treating a key as a parameter in the param dictionary may work, then .dropping it with an adapter so it isn’t learned.
