google-deepmind / google-deepmind/chex

Combining assert_scalar_positive with jax.vmap

Open
#389 1 comment 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
957
Forks
74
Avg merge
21h 10m
Merged PRs (30d)
1

Description

Is it possible to apply `assert_scalar_positive` to a vector? The MWE below shows what I'd like to do; however, I cannot currently see how this is possible with Chex.

```py
import jax.numpy as jnp
import jax
from chex import assert_scalar_positive

x_scaler = 1.
x_vector = jnp.array([1.,1.])

assert_scalar_positive(x_scaler) # Works
jax.vmap(assert_scalar_positive)(x_vector) # What I'd like
```

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.