import jax import jaxlib !cat /var/colab/hostname print(jax.__version__) print(jaxlib.__version__) from jaxlib import xla_extension import jax key = jax.random.PRNGKey(1701) arr = jax.random.normal(key, (1000,)) device = arr.device_buffer.device() print(f"JAX device type: {device}") assert isinstance(device, xla_extension.GpuDevice), "unexpected JAX device type" import jax import numpy as np # matrix multiplication on GPU key = jax.random.PRNGKey(0) x = jax.random.normal(key, (3000, 3000)) result = jax.numpy.dot(x, x.T).mean() print(result) import jax.numpy as jnp import jax.random as rand N = 10 M = 20 key = rand.PRNGKey(1701) X = rand.normal(key, (N, M)) u, s, vt = jnp.linalg.svd(X) assert u.shape == (N, N) assert vt.shape == (M, M) print(s) @jax.jit def selu(x, alpha=1.67, lmbda=1.05): return lmbda * jax.numpy.where(x > 0, x, alpha * jax.numpy.exp(x) - alpha) x = jax.random.normal(key, (5000,)) result = selu(x).block_until_ready() print(result)