(fn)
| 649 | key, child_key = jax.random.split(key) |
| 650 | base_results.append(fn(child_key)) |
| 651 | base_results = jnp.stack(base_results) |
| 652 | |
| 653 | def batch(fn): |
| 654 | return lambda split_keys: jax.vmap(fn)(split_keys.keys) |
| 655 |
no outgoing calls
no test coverage detected