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