Files
hpos-data/train/tpu.py
Pritimay Sarkar 35479623f7 tpu
2024-03-06 23:33:52 +05:30

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)