Home Knowledge Base Google TPU (Tensor Processing Unit)

Google TPU (Tensor Processing Unit)

What is TPU? Purpose-built ASIC for ML training and inference, available via Google Cloud.

TPU Versions

VersionYearFeatures
TPU v22017180 TFLOPS
TPU v32018420 TFLOPS, liquid cooled
TPU v42021275 TFLOPS, 4096-chip pods
TPU v5e2023Cost-optimized inference
TPU v5p2023Training-optimized

TPU Architecture

Using TPUs with JAX

import jax
import jax.numpy as jnp

# Check TPU availability
print(jax.devices())  # [TpuDevice(...)]

# Arrays automatically use TPU
x = jnp.ones((1000, 1000))
y = jnp.dot(x, x)

Multi-TPU Training

from jax.sharding import PartitionSpec, NamedSharding
from jax.experimental import mesh_utils

# Create device mesh
devices = mesh_utils.create_device_mesh((4, 2))  # 4x2 TPU grid

# Shard data across devices
mesh = Mesh(devices, axis_names=("data", "model"))
sharding = NamedSharding(mesh, PartitionSpec("data", None))

# Distribute array
distributed_data = jax.device_put(data, sharding)

TPU with TensorFlow

import tensorflow as tf

# TPU initialization
resolver = tf.distribute.cluster_resolver.TPUClusterResolver()
tf.config.experimental_connect_to_cluster(resolver)
tf.tpu.experimental.initialize_tpu_system(resolver)

# Create strategy
strategy = tf.distribute.TPUStrategy(resolver)

with strategy.scope():
    model = create_model()
    model.compile(...)
model.fit(dataset)

TPU vs GPU Comparison

AspectTPUGPU (H100)
Best forGoogle ecosystemGeneral
Memory16-64GB HBM80GB HBM
InterconnectTPU podsNVLink
SoftwareJAX/TFPyTorch/TF
AvailabilityGCP onlyUniversal

Pricing (GCP)

TypeOn-demandSpot
TPU v4$3.22/hr$0.97/hr
TPU v5e$1.20/hr$0.36/hr

Best Practices

tpugoogletensor

Explore 500+ Semiconductor & AI Topics

From EUV lithography to CUDA optimization — search the full knowledge base or chat with our AI assistant.