This commit is contained in:
Pritimay Sarkar
2024-03-06 23:33:52 +05:30
parent 482f0486c0
commit 35479623f7

19
train/tpu.py Normal file
View 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)