20 lines
390 B
Python
20 lines
390 B
Python
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)
|