之前一直在 GPU 集群上跑大规模模型训练,直到项目迁移到 Google Cloud TPU 时才发现,仅仅把训练框架换成 TPU 版本是远远不够的。整个编译链路、数据管道、算子融合策略和显存规划都变了,踩了一圈坑之后,才慢慢把 TPU 软件栈的运作方式理清楚。这篇文章会把 Google TPU 软件栈的核心组件拆开梳理,并结合 JAX、TensorFlow 分布式策略给出可落地的训练示例。适合正在评估 TPU、刚拿到 TPU 配额准备迁移训练任务、或者已经遇到底层报错但搜不到完整解释的开发者。读完你能掌握 TPU 的完整软件链路、训练启动方式、常见编译报错定位思路,以及把模型真正跑在 TPU 上而不只是跑在“模拟器里”的实践经验。
1. TPU 到底是什么:不止是“GPU 的替代品”
1.1 AI 时代对芯片提出的新要求
过去十年,深度学习算力的主力是 GPU。GPU 的核心优势在于并行计算能力很强,可以同时处理大量简单运算,非常适合矩阵乘法和卷积这类深度学习算子。但随着模型规模快速膨胀,尤其是大规模 Transformer、推荐系统和多模态模型出现之后,算力需求不再只看单卡峰值,还要看集群扩展效率、内存带宽、编译优化程度和单位算力的性价比。
Google 正是在这种背景下推出了 TPU(Tensor Processing Unit,张量处理单元)。它不是为了替代 GPU 而设计的通用芯片,而是为了加速 TensorFlow 和 JAX 这类深度学习框架中的张量运算而专门定制的 ASIC 芯片。你可以把它理解成一个“为矩阵乘法而生的专用加速器”,在特定负载下的能效比和吞吐表现非常突出。
这里的“专用”意味着两件事:一方面,TPU 在矩阵运算、卷积运算和大规模分布式训练上性能很强;另一方面,它对模型结构的支持并不是无条件的,某些自定义算子如果不适配 XLA 编译器,会直接卡在编译阶段。
1.2 CPU、GPU、TPU、NPU 的定位区别
先用一个表格把这几种芯片的定位区分清楚:
| 芯片 | 全称 | 定位 | 典型场景 |
|---|---|---|---|
| CPU | Central Processing Unit | 通用计算,强顺序执行 | 操作系统、数据库、逻辑控制 |
| GPU | Graphics Processing Unit | 并行计算,通用加速 | 深度学习训练、渲染、科学计算 |
| TPU | Tensor Processing Unit | 张量专用 ASIC | TensorFlow/JAX 大规模训练与推理 |
| NPU | Neural Processing Unit | 神经网络专用处理器 | 手机端推理、边缘 AI、端侧加速 |
从开发者的角度看,这套术语体系经常混用,尤其是在端侧场景里 NPU 和 TPU 的界限并不明显。但如果你做的是云上大规模训练,TPU 和 GPU 的差异会直接影响代码写法、编译器行为、数据加载策略和故障排查方式。
1.3 TPU 软件栈的整体构成
TPU 硬件本身只是一块芯片,真正让它跑起来的是软件栈。Google TPU 软件栈通常可以分成几层:
- 前端框架层:提供 TensorFlow、JAX、PyTorch 等框架的 TPU 后端接口。
- 编译器层:XLA(Accelerated Linear Algebra)编译器承担了关键的角色。
- 运行时层:libtpu、TPU Runtime、设备驱动等负责与硬件通信。
- 系统调度层:在 Cloud TPU 场景下负责资源创建、调度和生命周期管理。
下面这张 ASCII 简图可以帮助理解数据流向:
PyTorch / TensorFlow / JAX ↓ 前端图表示(Graph / Program) ↓ XLA 编译器(HLO → 优化 → LLVM / TPU 指令) ↓ TPU Runtime(libtpu / 设备驱动) ↓ TPU 硬件理解这个软件栈是很有必要的,因为后续很多报错都不是框架层报出来的,而是 XLA 编译阶段抛出的。你不能只盯着 TensorFlow 的报错堆栈看,还得顺着 XLA 的报错信息往下追。
2. TPU 软件栈核心组件拆解
2.1 XLA 编译器:从框架图到 TPU 指令的桥梁
XLA 是 TPU 软件栈里最核心、最容易被忽视的一层。XLA 的全称是 Accelerated Linear Algebra,它是 Google 推出的领域专用编译器,专门用来把 TensorFlow、JAX 等框架的计算图编译成高效的底层指令。
XLA 的工作过程大致如下:
- 框架层把计算任务表达成计算图或程序。
- XLA 将计算图转换成 HLO(High Level Operations,高级操作)表示。
- XLA 在 HLO 层面上做优化,包括算子融合(Fusion)、常量折叠(Constant Folding)、内存分配规划、并行化调度等。
- 最终把 HLO 编译成目标设备的 LLVM IR 或 TPU 专用指令序列。
在 JAX 中,@jax.jit装饰器就是触发 XLA 编译的入口。当你给一个函数加上@jax.jit时,JAX 会做两件事:把函数转换成计算图,然后交给 XLA 编译。
import jax import jax.numpy as jnp # 加 jit 后,函数不会一行一行执行,而是整体编译后执行 @jax.jit def linear(x, w, b): return jnp.dot(x, w) + b x = jnp.ones((8, 128)) w = jnp.ones((128, 64)) b = jnp.ones((64,)) y = linear(x, w, b) print(y.shape)这里需要注意,@jax.jit并不是简单地把函数内部代码“合并成一个整体”,而是让 XLA 拿到整个计算流程后做算子融合。比如上面的matmul + add,XLA 在编译后很可能融合成一个融合算子(Fusion),减少核函数启动次数和设备缓存回写的次数。
生产环境里,XLA 编译失败通常表现为类似“Detected unsupported operations when trying to compile graph”的报错,这类问题不是简单的语法错误,而是某个算子 XLA 还不支持或无法高效融合。
2.2 JAX 与 TensorFlow 对 TPU 的支持方式
JAX 是目前 Google 官方推荐的 TPU 训练框架之一。JAX 的设计思路是把 NumPy 风格的 API 和自动微分、JIT 编译结合,再配合pmap、shard_map等并行抽象,天然适合 TPU 这种需要精确控制数据切分的加速器。
TensorFlow 对 TPU 的支持则主要通过TPUStrategy实现。TPUStrategy是 TensorFlow 分布式策略中的一种,它负责把模型变量、优化器状态和数据分配到多个 TPU 核心上。
PyTorch 用户也不用太担心,PyTorch/XLA 项目已经提供了torch_xla包,让 PyTorch 模型可以跑在 TPU 上,同时支持 XLA 编译优化。不过平心而论,PyTorch 在 TPU 上的生态成熟度目前不如 JAX 和 TensorFlow,如果你是从 PyTorch 社区迁移过来的,建议预留更多时间做算子兼容性验证。
2.3 libtpu 与低层运行时
libtpu 是 TPU 的低层运行时库,负责主机与 TPU 设备之间的通信、内存管理、指令提交等。在 Cloud TPU 环境中,用户代码通过 gRPC 与 TPU Worker 通信,但具体指令的下发是由 libtpu 完成的。
这里想强调一个容易踩坑的点:TPU 的虚拟地址页大小是 16KB,而绝大多数 x86 Linux 主机默认是 4KB 页大小。当你用 pip 安装编译好的包时,如果包本身是在 4KB 页环境下编译的,运行到低层库调用阶段可能出现页面大小不匹配的报错,这就是网上常见的“an error occurred while preparing sdk package 16 kb page size”一类问题的底层背景。
解决思路通常不是修改内核页大小,而是从官方渠道获取与 TPU 运行环境匹配的预编译包,或者在你的 TPU 虚拟机上重新编译相关依赖,而不是直接把本地 x86 环境里的包复制到 TPU 环境。
3. 环境准备与 TPU 获取方式
3.1 Cloud TPU 与 Colab 的选择
搭建 TPU 软件栈的第一步是获取 TPU 环境。根据场景不同,有几种主流选择:
- Google Cloud TPU:适合长期训练任务,支持创建 TPU Pod 和 TPU VM。
- Google Colab:提供免费的 TPU 运行时,适合学习和小规模实验。
- TPU Research Cloud:面向研究者的免费 TPU 配额项目,适合科研场景。
- 本地模拟器:适合调试代码逻辑,但性能完全不能代替真实 TPU。
如果你是第一次接触 TPU,建议从 Colab 开始,它能用很短的时间验证你的代码是否走通了 TPU 软件链路。等代码稳定后再上 Cloud TPU 集群做真实规模训练。
下面通过一个简单的命令检查 TPU 运行时环境关键信息:
# 确认 Python 版本 python3 --version # 确认 TPU 设备是否可见(在 TPU VM 或 Colab TPU 上执行) ls /dev/accel* 2>/dev/null || echo "no accelerator device found" # 查看安装的关键包版本 pip list 2>/dev/null | grep -Ei "jax|tensorflow|torch|xla"在实际 Cloud TPU v4 或更高版本的 TPU VM 环境中,设备节点通常不是/dev/tpu而是/dev/accel0、/dev/accel1这类路径,这一点和传统 GPU 环境有所不同。
3.2 JAX 环境变量与初始化
JAX 在 TPU 上运行前,需要确认设备已经初始化成功。在 Colab 的较新版本 JAX 中,通常会自动识别 TPU,但如果你使用的是旧版本,或使用自定义镜像,就需要手动初始化:
import jax import jax.numpy as jnp # 有些旧版本需要手动初始化 TPU # 新版 JAX 一般会自动初始化,无需显式调用 try: from jax.tools import colab_tpu colab_tpu.setup_tpu() except Exception as e: print("Manual TPU init skipped or already supported:", e) # 查看当前所有可用设备 devices = jax.devices() print("JAX devices:", devices) print("JAX default backend:", jax.default_backend())运行正常的情况下,jax.devices()会返回 TPU 设备列表,jax.default_backend()返回tpu。如果你的输出是cpu或gpu,说明 JAX 并没有真正访问 TPU,需要检查镜像版本或运行环境。
有一个经常被忽略的问题:Colab TPU 切换后,必须重启运行时、重新安装 JAX 版本,否则 JAX 内部缓存的设备信息不会刷新。如果切换 TPU 类型后jax.devices()仍然显示旧的设备列表,优先考虑重启运行时而不是调试代码。
3.3 从源代码编译还是预编译包
在 TPU 上安装 JAX 时,我一直建议优先使用官方发布的预编译包,因为这些包会和 TPU 的页大小、指令集和运行时库做过匹配测试。只有在需要修改 JAX 源码、debug 底层行为,或者在 TPU VM 上做二次开发时,才考虑从源码编译。
下面是一个在 TPU VM 上安装 JAX 的典型流程:
# 激活虚拟环境 python3 -m venv venv source venv/bin/activate # 安装 JAX 的 TPU 版本 # 更准确的安装命令需要参考官方文档,关键是根据 TPU 版本和系统架构选择匹配的 wheel pip install --upgrade "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html # 验证安装 python3 -c "import jax; print(jax.devices())"需要注意,jax[tpu]这个 extra 的依赖列表和可用的 wheel 索引会随着版本变化,建议以官方 JAX 仓库或 Google Cloud 文档中的安装命令为准。不匹配的版本经常导致 libtpu 无法加载,报错信息里会出现类似“Could not load libtpu.so”的字样。
4. 完整实战:用 JAX 在 TPU 上训练一个图像分类模型
讲完概念和准备,下面进入完整实战环节。这里以 JAX 为例,一步步演示从数据加载到 TPU 训练的全过程。
4.1 创建项目结构与准备数据
先创建一个项目目录,把代码按模块划分好:
tpu-jax-demo/ ├── main.py ├── model.py ├── data.py ├── train.py └── requirements.txt数据部分使用 TensorFlow Datasets 中的 MNIST 数据集。MNIST 虽然简单,但它覆盖了数据加载、批量切分、训练循环、模型保存的完整流程,很适合用来验证 TPU 软件栈是否正常工作。
创建requirements.txt:
jax[tpu] tensorflow-cpu tensorflow-datasets optax这里特意引入optax作为优化器库,它是 JAX 生态中最常用的优化器集合,后续如果要替换成 AdamW、LAMB 或自定义学习率调度,直接改一行配置即可。
4.2 编写数据加载模块
创建data.py:
import tensorflow_datasets as tfds # 加载 MNIST 数据并转换为 NumPy 数组 def load_mnist(batch_size=128): ds = tfds.load("mnist", split=["train", "test"], as_supervised=True) def prepare(dataset): dataset = dataset.map(lambda x, y: (tf.cast(x, tf.float32) / 255.0, y)) dataset = dataset.batch(batch_size) dataset = dataset.prefetch(tf.data.AUTOTUNE) return dataset train_ds = prepare(ds[0]) test_ds = prepare(ds[1]) return train_ds, test_ds在 TPU 训练中,数据加载是一个经常被忽略的瓶颈。CPU 端的数据加载速度如果跟不上 TPU 的消费速度,训练曲线会出现明显的周期性停顿。prefetch是解决这个问题的最简单手段。
4.3 定义模型与训练循环
创建model.py:
import jax.numpy as jnp from flax import linen as nn # 简单卷积网络 class SimpleCNN(nn.Module): @nn.compact def __call__(self, x, training: bool = True): x = nn.Conv(features=32, kernel_size=(3, 3), padding="SAME")(x) x = nn.relu(x) x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2)) x = nn.Conv(features=64, kernel_size=(3, 3), padding="SAME")(x) x = nn.relu(x) x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2)) x = x.reshape((x.shape[0], -1)) x = nn.Dense(features=128)(x) x = nn.relu(x) x = nn.Dense(features=10)(x) return x这里使用 Flax 来定义模型。Flax 是 Google 官方维护的 JAX 神经网络库,和 JAX 一起使用时体验最自然。如果你之前用过 PyTorch 的nn.Module,Flax 的nn.Module风格会比较接近。
创建train.py,这是训练循环的核心:
import jax import jax.numpy as jnp import optax from flax.training import train_state from model import SimpleCNN from data import load_mnist def cross_entropy_loss(logits, labels): one_hot = jax.nn.one_hot(labels, num_classes=10) return -jnp.mean(jnp.sum(one_hot * jax.nn.log_softmax(logits), axis=-1)) def compute_metrics(logits, labels): loss = cross_entropy_loss(logits, labels) accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == labels) return loss, accuracy @jax.jit def train_step(state, batch): images, labels = batch def loss_fn(params): logits = state.apply_fn({"params": params}, images, training=True) return cross_entropy_loss(logits, labels) loss, grads = jax.value_and_grad(loss_fn)(state.params) state = state.apply_gradients(grads=grads) return state, loss @jax.jit def eval_step(state, batch): images, labels = batch logits = state.apply_fn({"params": state.params}, images, training=False) return compute_metrics(logits, labels) def create_train_state(rng, learning_rate): model = SimpleCNN() params = model.init(rng, jnp.ones((1, 28, 28, 1)))["params"] tx = optax.adam(learning_rate) return train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx) def main(): rng = jax.random.PRNGKey(0) state = create_train_state(rng, learning_rate=1e-3) train_ds, test_ds = load_mnist(batch_size=128) # 将 TensorFlow Dataset 转换为 NumPy Iterator train_iter = iter(train_ds) test_iter = iter(test_ds) for epoch in range(3): for step in range(100): batch = next(train_iter) images = batch[0].numpy() labels = batch[1].numpy() state, loss = train_step(state, (images, labels)) if step % 20 == 0: print(f"epoch {epoch} step {step} loss {loss:.4f}") # 每个 epoch 结束评估一次 total_loss = 0.0 total_acc = 0.0 num_batches = 0 for _ in range(50): batch = next(test_iter) images = batch[0].numpy() labels = batch[1].numpy() loss, acc = eval_step(state, (images, labels)) total_loss += loss total_acc += acc num_batches += 1 print(f"epoch {epoch} eval loss {total_loss / num_batches:.4f} " f"acc {total_acc / num_batches:.4f}") if __name__ == "__main__": main()这段代码有几个关键点:
- 使用
@jax.jit装饰训练和评估函数,让 XLA 将整个计算过程编译成融合算子。 - 每个 batch 通过
.numpy()从 TensorFlow Dataset 转换为 NumPy 数组,供 JAX 消费。 - 训练状态由
TrainState统一管理,包含模型参数和优化器状态。 - 打印网络在 MNIST 上的 loss 和 acc,用来验证模型真实地训练起来了。
4.4 真机运行与验证
在 Colab 或 TPU VM 上执行:
python3 train.py正常输出会类似:
epoch 0 step 0 loss 2.3021 epoch 0 step 20 loss 0.4218 epoch 0 step 40 loss 0.2534 epoch 0 step 60 loss 0.1842 epoch 0 step 80 loss 0.1507 epoch 0 eval loss 0.0852 acc 0.9734 ...看到 loss 在下降、eval accuracy 稳步上升,说明 JAX + XLA + TPU 这条链路已经跑通了。接下来就可以把这里的SimpleCNN替换成真实模型,把load_mnist替换成你的真实数据集。
4.5 关于多卡 TPU 的扩展思路
上面的代码是单进程、单 TPU 核心训练的写法。如果你创建的是多核心 TPU,JAX 会自动识别多个设备。想要利用多个核心并行训练,最简单的方式是使用jax.device_put和pmap对 batch 做数据并行切分。
from jax import pmap # 将状态复制到所有设备 state = jax.device_put_replicated(state, jax.devices()) # 多设备并行训练步骤 @pmap def train_step_multi(state, batch): return train_step(state, batch) # 切分 batch 到多个设备 images = images.reshape((num_devices, -1) + images.shape[1:]) labels = labels.reshape((num_devices, -1))这只是一个非常简化的pmap示例,真实多卡训练还要考虑 batch 切分策略、梯度累积、AllReduce 等细节。建议先把单 core 代码跑通,再逐步过渡到pmap或shard_map。
5. TensorFlow 方式:用 TPUStrategy 做分布式训练
虽然 JAX 在 TPU 上体验越来越主流,但生产环境中大量存量代码还是 TensorFlow 的。如果你不想重写模型,可以直接使用 TensorFlow 的TPUStrategy把现有模型迁移到 TPU。
5.1 TPUStrategy 工作原理
TPUStrategy是 TensorFlow 的分布式策略之一。它会自动完成几件事:
- 把模型变量复制到每个 TPU 核心。
- 把全局 batch 切分成每个核心处理一个子 batch。
- 优化器在反向传播后做跨核心梯度 AllReduce。
- 对训练循环内部使用
tf.function编译,配合 XLA 加速。
它的关键思路是“数据并行 + 同步更新”。用户只需要写一份单机单卡代码,策略层负责把计算扩展到多个 TPU 核心上。
5.2 TPUStrategy 训练示例
import tensorflow as tf # 1. 初始化 TPU resolver = tf.distribute.cluster_resolver.TPUClusterResolver() tf.config.experimental_connect_to_cluster(resolver) tf.tpu.experimental.initialize_tpu_system(resolver) strategy = tf.distribute.TPUStrategy(resolver) print("TPU devices:", resolver.cluster_spec().as_dict()) # 2. 在 strategy 作用域内构建模型 with strategy.scope(): model = tf.keras.Sequential([ tf.keras.layers.Conv2D(32, (3, 3), activation="relu", input_shape=(28, 28, 1)), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.Conv2D(64, (3, 3), activation="relu"), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.Flatten(), tf.keras.layers.Dense(128, activation="relu"), tf.keras.layers.Dense(10), ]) model.compile( optimizer=tf.keras.optimizers.Adam(1e-3), loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=["accuracy"], ) # 3. 准备数据 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() x_train = x_train.reshape((-1, 28, 28, 1)).astype("float32") / 255.0 x_test = x_test.reshape((-1, 28, 28, 1)).astype("float32") / 255.0 # 4. 训练 model.fit( x_train, y_train, batch_size=128, epochs=3, validation_data=(x_test, y_test), )TPUStrategy的迁移成本很低,尤其适合原本就是 Keras 写的模型。如果你的模型里含有自定义tf.Variable更新逻辑,就需要确认变量创建是否都发生在strategy.scope()内。
5.3 TFRecord 数据管道与 TPU 的配合
真实训练中数据通常在 GCS 上,推荐使用 TFRecord 格式来减少小文件数量,提高 TPU 的读取效率。
import tensorflow as tf def decode_fn(record_bytes): features = { "image": tf.io.FixedLenFeature([28 * 28], tf.float32), "label": tf.io.FixedLenFeature([], tf.int64), } parsed = tf.io.parse_single_example(record_bytes, features) image = tf.reshape(parsed["image"], (28, 28, 1)) return image, parsed["label"] def build_dataset(file_pattern, batch_size, is_training=True): dataset = tf.data.Dataset.list_files(file_pattern) dataset = dataset.interleave( tf.data.TFRecordDataset, cycle_length=8, num_parallel_calls=tf.data.AUTOTUNE, ) dataset = dataset.map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE) if is_training: dataset = dataset.shuffle(10000).repeat() dataset = dataset.batch(batch_size, drop_remainder=True) dataset = dataset.prefetch(tf.data.AUTOTUNE) return dataset数据管道本身是 CPU 上的工作,但它的吞吐能力直接决定了 TPU 是否被喂饱。经验法则是:数据管道的吞吐至少要达到模型训练速度的 2 倍,否则训练过程中会不断出现“等待数据”的空窗期。
6. 常见问题与排查思路
6.1 TPU 初始化失败或设备不可见
问题现象:
RuntimeError: Unable to initialize the TPU system.可能原因:
- 当前运行时不是 TPU 环境。
- JAX/TensorFlow 版本与 TPU 运行时版本不匹配。
- 运行时使用旧内核,没有加载必要的驱动。
排查步骤:
- 先检查设备文件是否存在,在 TPU VM 上执行
ls /dev/accel*。 - 再执行
jax.devices()或tf.config.list_logical_devices("TPU")看框架层是否识别到设备。 - 确认虚拟环境中的包版本与当前 TPU 版本匹配。
- 重启运行时,尤其是在切换 TPU 类型之后。
6.2 XLA 编译报错 “unsupported operations”
问题现象:
Detected unsupported operations when trying to compile graph可能原因:
- 模型里使用了 XLA 不支持的算子。
- 自定义 Layer 或自定义 JAX 函数里调用了无法被转换成 HLO 的 Python 控制流。
解决思路:
- 逐步裁剪模型,定位到具体是哪个操作不能被编译。
- 检查是否是动态形状问题,TPU 上尽量避免 runtime shape 变化。
- 对自定义算子优先考虑能否用现有的 JAX 原生操作重写。
6.3 关于 16KB page size 编译包报错
问题现象:
在安装或运行 TPU SDK 相关包时,出现类似 “an error occurred while preparing sdk package 16 kb page size” 的报错。
背景原因:
TPU 运行环境的虚拟地址页大小是 16KB,而常见 x86 Linux 是 4KB 页。本地下载的预编译 wheel 如果是在 4KB 页环境下构建的,运行时会与 TPU 内核模块或运行时库不兼容。
解决思路:
- 不要直接从普通 PyPI 镜像安装 TPU 相关包,优先使用官方发布渠道。
- 在 TPU VM 内重新安装匹配版本,而不是把本地环境复制过去。
- 查看完整日志,确认是下载失败还是安装后运行时崩溃。
6.4 OOM 与 batch size 调整
问题现象:
训练时出现Resource exhausted错误。
可能原因:
- batch size 过大,超出了 TPU 单核的 HBM 容量。
- 模型参数或中间激活值过大。
- 数据管道
prefetch使用内存过多。
解决思路:
- 先尝试减小 batch size,确认模型在单卡上可以跑通。
- 合理设置
prefetch和num_parallel_calls,避免 CPU 端内存溢出。 - 检查 XLA 是否启用了内存优化选项,但不要盲目相信默认配置。
6.5 排查清单汇总
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 设备不可见 | 运行时环境错误 / 版本不匹配 | 检查/dev/accel*,重启运行时,核对版本 |
| XLA 编译失败 | 不支持的算子 / 动态 shape | 简化模型逐步定位,避免动态形状 |
| 16KB page size 报错 | wheel 与 TPU 页大小不匹配 | 使用官方发布渠道安装匹配包 |
| OOM | batch size 过大 / 中间激活过大 | 调整 batch size,优化数据管道内存 |
| 训练速度上不去 | 数据管道存在瓶颈 | 使用 TFRecord、interleave、prefetch 优化 |
| 输出结果有 NaN | 学习率过高 / 精度配置不当 | 降低学习率,使用混合精度时检查损失缩放 |
7. 最佳实践与工程建议
7.1 数据管道要提前压测
TPU 非常“挑食”,它的运算速度快到如果数据管道没跟上,训练就会空等。建议在正式训练前,先单独测试数据管道的吞吐能力。
import time train_iter = iter(train_ds) start = time.time() for i in range(100): batch = next(train_iter) end = time.time() print(f"100 batches take {end - start:.2f}s")如果数据读取时间占比过高,优先使用 TFRecord 格式、interleave并行读取、prefetch预加载等手段。
7.2 使用混合精度时需要验证损失缩放
TPU 的 bfloat16 支持是它的强项之一。用jax时可以很自然地把部分参数转换成bfloat16或使用混合精度训练。但在混合精度下,梯度很小的时候有可能在低精度下溢出,导致 NaN。
建议开启损失缩放(Loss Scaling)机制,并周期性检查梯度统计,而不是在出现 NaN 后才去排查。
7.3 模型算子优先考虑 JAX/TensorFlow 原生实现
遇到自定义算子时,优先检查原生库里有没有替代实现。XLA 对原生算子的融合优化做得很成熟,但自定义算子往往无法被融合,会打乱 XLA 的优化策略。如果一定要使用自定义算子,至少把“无法编译”的算子独立出来,避免拖累整体性能。
7.4 成本与资源管理
TPU 资源通常是按时计费的,长期任务建议配合 checkpoint 和自动重启机制。这里给出几个实用性建议:
- 训练脚本要有稳定的 checkpoint 保存与恢复逻辑。
- 使用抢占式资源时,代码要能安全处理进程中途被回收的情况。
- 创建 TPU 实例后尽快验证训练脚本,避免空转计费。
- 配置监控报警,对异常掉线、训练停滞、指标不回传这些情况做告警。
7.5 把训练过程日志化
TPU 上的日志查看比 GPU 环境要麻烦一些,尤其是多 worker 场景下。推荐在训练脚本中显式记录关键指标,比如每个 epoch 的 loss、acc、每个 step 的平均耗时、数据加载耗时等,并把日志输出到统一的日志平台,方便后续归因。
import logging logging.basicConfig( level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s", ) logging.info("start training, devices=%s", jax.devices())这一步可能看起来不起眼,但当你遇到 TPU 训练中途卡死、重启、性能退化问题时,这些日志是定位问题的唯一线索。
8. 总结
Google TPU 软件栈本质上是一条从深度学习框架到专用硬件的编译和运行链路,JAX、TensorFlow、XLA 和 libtpu 各司其职。对开发者而言,迁移到 TPU 时最需要调整的往往不是模型结构本身,而是对计算图编译、数据切分、内存规划这些底层层面的理解。
如果你准备开始尝试,建议从 Colab 免费 TPU 环境入手,把 MNIST 或自己的小型模型跑通,再逐步扩展到真实业务模型。遇到 16KB page size、XLA 编译失败、TPU 设备不可见这类问题,先顺着软件栈的层级逐层排查,不要一开始就怀疑硬件。只要把 JAX 或 TensorFlow 的 TPU 链路跑通一遍,后续在 Cloud TPU 上做规模化训练会顺畅得多。