tpu
This commit is contained in:
19
train/tpu.py
Normal file
19
train/tpu.py
Normal file
@@ -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)
|
||||
Reference in New Issue
Block a user