to_jax#

t3toolbox.backend.common.to_jax(x)#
def to_jax(x):

jnp.array(x) when jax is available; otherwise np.array(x) with a warning (see jax_or_warn).