diff --git a/train/tpu.py b/train/tpu.py new file mode 100644 index 0000000..86c771e --- /dev/null +++ b/train/tpu.py @@ -0,0 +1,19 @@ +import jax +import jax.numpy as jnp + +devices = jax.devices() +print("Available devices:", devices) + +tpu_device = devices[0] + +@jax.jit +def simple_function(x): + return jnp.sum(x ** 2) + +input_data = jax.random.uniform(jax.random.PRNGKey(0), shape=(1000, 1000)) + +input_data_on_tpu = jax.device_put(input_data, tpu_device) + +result = simple_function(input_data_on_tpu) + +print("Result:", result)