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).