戰(zhàn)解析)
之前一直在 GPU 集群上跑大規(guī)模模型訓(xùn)練直到項(xiàng)目遷移到 Google Cloud TPU 時(shí)才發(fā)現(xiàn)僅僅把訓(xùn)練框架換成 TPU 版本是遠(yuǎn)遠(yuǎn)不夠的。整個(gè)編譯鏈路、數(shù)據(jù)管道、算子融合策略和顯存規(guī)劃都變了踩了一圈坑之后才慢慢把 TPU 軟件棧的運(yùn)作方式理清楚。這篇文章會(huì)把 Google TPU 軟件棧的核心組件拆開梳理并結(jié)合 JAX、TensorFlow 分布式策略給出可落地的訓(xùn)練示例。適合正在評(píng)估 TPU、剛拿到 TPU 配額準(zhǔn)備遷移訓(xùn)練任務(wù)、或者已經(jīng)遇到底層報(bào)錯(cuò)但搜不到完整解釋的開發(fā)者。讀完你能掌握 TPU 的完整軟件鏈路、訓(xùn)練啟動(dòng)方式、常見編譯報(bào)錯(cuò)定位思路以及把模型真正跑在 TPU 上而不只是跑在“模擬器里”的實(shí)踐經(jīng)驗(yàn)。1. TPU 到底是什么不止是“GPU 的替代品”1.1 AI 時(shí)代對(duì)芯片提出的新要求過去十年深度學(xué)習(xí)算力的主力是 GPU。GPU 的核心優(yōu)勢(shì)在于并行計(jì)算能力很強(qiáng)可以同時(shí)處理大量簡(jiǎn)單運(yùn)算非常適合矩陣乘法和卷積這類深度學(xué)習(xí)算子。但隨著模型規(guī)??焖倥蛎浻绕涫谴笠?guī)模 Transformer、推薦系統(tǒng)和多模態(tài)模型出現(xiàn)之后算力需求不再只看單卡峰值還要看集群擴(kuò)展效率、內(nèi)存帶寬、編譯優(yōu)化程度和單位算力的性價(jià)比。Google 正是在這種背景下推出了 TPUTensor Processing Unit張量處理單元。它不是為了替代 GPU 而設(shè)計(jì)的通用芯片而是為了加速 TensorFlow 和 JAX 這類深度學(xué)習(xí)框架中的張量運(yùn)算而專門定制的 ASIC 芯片。你可以把它理解成一個(gè)“為矩陣乘法而生的專用加速器”在特定負(fù)載下的能效比和吞吐表現(xiàn)非常突出。這里的“專用”意味著兩件事一方面TPU 在矩陣運(yùn)算、卷積運(yùn)算和大規(guī)模分布式訓(xùn)練上性能很強(qiáng)另一方面它對(duì)模型結(jié)構(gòu)的支持并不是無條件的某些自定義算子如果不適配 XLA 編譯器會(huì)直接卡在編譯階段。1.2 CPU、GPU、TPU、NPU 的定位區(qū)別先用一個(gè)表格把這幾種芯片的定位區(qū)分清楚芯片全稱定位典型場(chǎng)景CPUCentral Processing Unit通用計(jì)算強(qiáng)順序執(zhí)行操作系統(tǒng)、數(shù)據(jù)庫、邏輯控制GPUGraphics Processing Unit并行計(jì)算通用加速深度學(xué)習(xí)訓(xùn)練、渲染、科學(xué)計(jì)算TPUTensor Processing Unit張量專用 ASICTensorFlow/JAX 大規(guī)模訓(xùn)練與推理NPUNeural Processing Unit神經(jīng)網(wǎng)絡(luò)專用處理器手機(jī)端推理、邊緣 AI、端側(cè)加速從開發(fā)者的角度看這套術(shù)語體系經(jīng)?;煊糜绕涫窃诙藗?cè)場(chǎng)景里 NPU 和 TPU 的界限并不明顯。但如果你做的是云上大規(guī)模訓(xùn)練TPU 和 GPU 的差異會(huì)直接影響代碼寫法、編譯器行為、數(shù)據(jù)加載策略和故障排查方式。1.3 TPU 軟件棧的整體構(gòu)成TPU 硬件本身只是一塊芯片真正讓它跑起來的是軟件棧。Google TPU 軟件棧通常可以分成幾層前端框架層提供 TensorFlow、JAX、PyTorch 等框架的 TPU 后端接口。編譯器層XLAAccelerated Linear Algebra編譯器承擔(dān)了關(guān)鍵的角色。運(yùn)行時(shí)層libtpu、TPU Runtime、設(shè)備驅(qū)動(dòng)等負(fù)責(zé)與硬件通信。系統(tǒng)調(diào)度層在 Cloud TPU 場(chǎng)景下負(fù)責(zé)資源創(chuàng)建、調(diào)度和生命周期管理。下面這張 ASCII 簡(jiǎn)圖可以幫助理解數(shù)據(jù)流向PyTorch / TensorFlow / JAX ↓ 前端圖表示Graph / Program ↓ XLA 編譯器HLO → 優(yōu)化 → LLVM / TPU 指令 ↓ TPU Runtimelibtpu / 設(shè)備驅(qū)動(dòng) ↓ TPU 硬件理解這個(gè)軟件棧是很有必要的因?yàn)楹罄m(xù)很多報(bào)錯(cuò)都不是框架層報(bào)出來的而是 XLA 編譯階段拋出的。你不能只盯著 TensorFlow 的報(bào)錯(cuò)堆??催€得順著 XLA 的報(bào)錯(cuò)信息往下追。2. TPU 軟件棧核心組件拆解2.1 XLA 編譯器從框架圖到 TPU 指令的橋梁XLA 是 TPU 軟件棧里最核心、最容易被忽視的一層。XLA 的全稱是 Accelerated Linear Algebra它是 Google 推出的領(lǐng)域?qū)S镁幾g器專門用來把 TensorFlow、JAX 等框架的計(jì)算圖編譯成高效的底層指令。XLA 的工作過程大致如下框架層把計(jì)算任務(wù)表達(dá)成計(jì)算圖或程序。XLA 將計(jì)算圖轉(zhuǎn)換成 HLOHigh Level Operations高級(jí)操作表示。XLA 在 HLO 層面上做優(yōu)化包括算子融合Fusion、常量折疊Constant Folding、內(nèi)存分配規(guī)劃、并行化調(diào)度等。最終把 HLO 編譯成目標(biāo)設(shè)備的 LLVM IR 或 TPU 專用指令序列。在 JAX 中jax.jit裝飾器就是觸發(fā) XLA 編譯的入口。當(dāng)你給一個(gè)函數(shù)加上jax.jit時(shí)JAX 會(huì)做兩件事把函數(shù)轉(zhuǎn)換成計(jì)算圖然后交給 XLA 編譯。import jax import jax.numpy as jnp # 加 jit 后函數(shù)不會(huì)一行一行執(zhí)行而是整體編譯后執(zhí)行 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并不是簡(jiǎn)單地把函數(shù)內(nèi)部代碼“合并成一個(gè)整體”而是讓 XLA 拿到整個(gè)計(jì)算流程后做算子融合。比如上面的matmul addXLA 在編譯后很可能融合成一個(gè)融合算子Fusion減少核函數(shù)啟動(dòng)次數(shù)和設(shè)備緩存回寫的次數(shù)。生產(chǎn)環(huán)境里XLA 編譯失敗通常表現(xiàn)為類似“Detected unsupported operations when trying to compile graph”的報(bào)錯(cuò)這類問題不是簡(jiǎn)單的語法錯(cuò)誤而是某個(gè)算子 XLA 還不支持或無法高效融合。2.2 JAX 與 TensorFlow 對(duì) TPU 的支持方式JAX 是目前 Google 官方推薦的 TPU 訓(xùn)練框架之一。JAX 的設(shè)計(jì)思路是把 NumPy 風(fēng)格的 API 和自動(dòng)微分、JIT 編譯結(jié)合再配合pmap、shard_map等并行抽象天然適合 TPU 這種需要精確控制數(shù)據(jù)切分的加速器。TensorFlow 對(duì) TPU 的支持則主要通過TPUStrategy實(shí)現(xiàn)。TPUStrategy是 TensorFlow 分布式策略中的一種它負(fù)責(zé)把模型變量、優(yōu)化器狀態(tài)和數(shù)據(jù)分配到多個(gè) TPU 核心上。PyTorch 用戶也不用太擔(dān)心PyTorch/XLA 項(xiàng)目已經(jīng)提供了torch_xla包讓 PyTorch 模型可以跑在 TPU 上同時(shí)支持 XLA 編譯優(yōu)化。不過平心而論P(yáng)yTorch 在 TPU 上的生態(tài)成熟度目前不如 JAX 和 TensorFlow如果你是從 PyTorch 社區(qū)遷移過來的建議預(yù)留更多時(shí)間做算子兼容性驗(yàn)證。2.3 libtpu 與低層運(yùn)行時(shí)libtpu 是 TPU 的低層運(yùn)行時(shí)庫負(fù)責(zé)主機(jī)與 TPU 設(shè)備之間的通信、內(nèi)存管理、指令提交等。在 Cloud TPU 環(huán)境中用戶代碼通過 gRPC 與 TPU Worker 通信但具體指令的下發(fā)是由 libtpu 完成的。這里想強(qiáng)調(diào)一個(gè)容易踩坑的點(diǎn)TPU 的虛擬地址頁大小是 16KB而絕大多數(shù) x86 Linux 主機(jī)默認(rèn)是 4KB 頁大小。當(dāng)你用 pip 安裝編譯好的包時(shí)如果包本身是在 4KB 頁環(huán)境下編譯的運(yùn)行到低層庫調(diào)用階段可能出現(xiàn)頁面大小不匹配的報(bào)錯(cuò)這就是網(wǎng)上常見的“an error occurred while preparing sdk package 16 kb page size”一類問題的底層背景。解決思路通常不是修改內(nèi)核頁大小而是從官方渠道獲取與 TPU 運(yùn)行環(huán)境匹配的預(yù)編譯包或者在你的 TPU 虛擬機(jī)上重新編譯相關(guān)依賴而不是直接把本地 x86 環(huán)境里的包復(fù)制到 TPU 環(huán)境。3. 環(huán)境準(zhǔn)備與 TPU 獲取方式3.1 Cloud TPU 與 Colab 的選擇搭建 TPU 軟件棧的第一步是獲取 TPU 環(huán)境。根據(jù)場(chǎng)景不同有幾種主流選擇Google Cloud TPU適合長(zhǎng)期訓(xùn)練任務(wù)支持創(chuàng)建 TPU Pod 和 TPU VM。Google Colab提供免費(fèi)的 TPU 運(yùn)行時(shí)適合學(xué)習(xí)和小規(guī)模實(shí)驗(yàn)。TPU Research Cloud面向研究者的免費(fèi) TPU 配額項(xiàng)目適合科研場(chǎng)景。本地模擬器適合調(diào)試代碼邏輯但性能完全不能代替真實(shí) TPU。如果你是第一次接觸 TPU建議從 Colab 開始它能用很短的時(shí)間驗(yàn)證你的代碼是否走通了 TPU 軟件鏈路。等代碼穩(wěn)定后再上 Cloud TPU 集群做真實(shí)規(guī)模訓(xùn)練。下面通過一個(gè)簡(jiǎn)單的命令檢查 TPU 運(yùn)行時(shí)環(huán)境關(guān)鍵信息# 確認(rèn) Python 版本 python3 --version # 確認(rèn) TPU 設(shè)備是否可見在 TPU VM 或 Colab TPU 上執(zhí)行 ls /dev/accel* 2/dev/null || echo no accelerator device found # 查看安裝的關(guān)鍵包版本 pip list 2/dev/null | grep -Ei jax|tensorflow|torch|xla在實(shí)際 Cloud TPU v4 或更高版本的 TPU VM 環(huán)境中設(shè)備節(jié)點(diǎn)通常不是/dev/tpu而是/dev/accel0、/dev/accel1這類路徑這一點(diǎn)和傳統(tǒng) GPU 環(huán)境有所不同。3.2 JAX 環(huán)境變量與初始化JAX 在 TPU 上運(yùn)行前需要確認(rèn)設(shè)備已經(jīng)初始化成功。在 Colab 的較新版本 JAX 中通常會(huì)自動(dòng)識(shí)別 TPU但如果你使用的是舊版本或使用自定義鏡像就需要手動(dòng)初始化import jax import jax.numpy as jnp # 有些舊版本需要手動(dòng)初始化 TPU # 新版 JAX 一般會(huì)自動(dòng)初始化無需顯式調(diào)用 try: from jax.tools import colab_tpu colab_tpu.setup_tpu() except Exception as e: print(Manual TPU init skipped or already supported:, e) # 查看當(dāng)前所有可用設(shè)備 devices jax.devices() print(JAX devices:, devices) print(JAX default backend:, jax.default_backend())運(yùn)行正常的情況下jax.devices()會(huì)返回 TPU 設(shè)備列表jax.default_backend()返回tpu。如果你的輸出是cpu或gpu說明 JAX 并沒有真正訪問 TPU需要檢查鏡像版本或運(yùn)行環(huán)境。有一個(gè)經(jīng)常被忽略的問題Colab TPU 切換后必須重啟運(yùn)行時(shí)、重新安裝 JAX 版本否則 JAX 內(nèi)部緩存的設(shè)備信息不會(huì)刷新。如果切換 TPU 類型后jax.devices()仍然顯示舊的設(shè)備列表優(yōu)先考慮重啟運(yùn)行時(shí)而不是調(diào)試代碼。3.3 從源代碼編譯還是預(yù)編譯包在 TPU 上安裝 JAX 時(shí)我一直建議優(yōu)先使用官方發(fā)布的預(yù)編譯包因?yàn)檫@些包會(huì)和 TPU 的頁大小、指令集和運(yùn)行時(shí)庫做過匹配測(cè)試。只有在需要修改 JAX 源碼、debug 底層行為或者在 TPU VM 上做二次開發(fā)時(shí)才考慮從源碼編譯。下面是一個(gè)在 TPU VM 上安裝 JAX 的典型流程# 激活虛擬環(huán)境 python3 -m venv venv source venv/bin/activate # 安裝 JAX 的 TPU 版本 # 更準(zhǔn)確的安裝命令需要參考官方文檔關(guān)鍵是根據(jù) TPU 版本和系統(tǒng)架構(gòu)選擇匹配的 wheel pip install --upgrade jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html # 驗(yàn)證安裝 python3 -c import jax; print(jax.devices())需要注意jax[tpu]這個(gè) extra 的依賴列表和可用的 wheel 索引會(huì)隨著版本變化建議以官方 JAX 倉庫或 Google Cloud 文檔中的安裝命令為準(zhǔn)。不匹配的版本經(jīng)常導(dǎo)致 libtpu 無法加載報(bào)錯(cuò)信息里會(huì)出現(xiàn)類似“Could not load libtpu.so”的字樣。4. 完整實(shí)戰(zhàn)用 JAX 在 TPU 上訓(xùn)練一個(gè)圖像分類模型講完概念和準(zhǔn)備下面進(jìn)入完整實(shí)戰(zhàn)環(huán)節(jié)。這里以 JAX 為例一步步演示從數(shù)據(jù)加載到 TPU 訓(xùn)練的全過程。4.1 創(chuàng)建項(xiàng)目結(jié)構(gòu)與準(zhǔn)備數(shù)據(jù)先創(chuàng)建一個(gè)項(xiàng)目目錄把代碼按模塊劃分好tpu-jax-demo/ ├── main.py ├── model.py ├── data.py ├── train.py └── requirements.txt數(shù)據(jù)部分使用 TensorFlow Datasets 中的 MNIST 數(shù)據(jù)集。MNIST 雖然簡(jiǎn)單但它覆蓋了數(shù)據(jù)加載、批量切分、訓(xùn)練循環(huán)、模型保存的完整流程很適合用來驗(yàn)證 TPU 軟件棧是否正常工作。創(chuàng)建requirements.txtjax[tpu] tensorflow-cpu tensorflow-datasets optax這里特意引入optax作為優(yōu)化器庫它是 JAX 生態(tài)中最常用的優(yōu)化器集合后續(xù)如果要替換成 AdamW、LAMB 或自定義學(xué)習(xí)率調(diào)度直接改一行配置即可。4.2 編寫數(shù)據(jù)加載模塊創(chuàng)建data.pyimport tensorflow_datasets as tfds # 加載 MNIST 數(shù)據(jù)并轉(zhuǎn)換為 NumPy 數(shù)組 def load_mnist(batch_size128): ds tfds.load(mnist, split[train, test], as_supervisedTrue) 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 訓(xùn)練中數(shù)據(jù)加載是一個(gè)經(jīng)常被忽略的瓶頸。CPU 端的數(shù)據(jù)加載速度如果跟不上 TPU 的消費(fèi)速度訓(xùn)練曲線會(huì)出現(xiàn)明顯的周期性停頓。prefetch是解決這個(gè)問題的最簡(jiǎn)單手段。4.3 定義模型與訓(xùn)練循環(huán)創(chuàng)建model.pyimport jax.numpy as jnp from flax import linen as nn # 簡(jiǎn)單卷積網(wǎng)絡(luò) class SimpleCNN(nn.Module): nn.compact def __call__(self, x, training: bool True): x nn.Conv(features32, kernel_size(3, 3), paddingSAME)(x) x nn.relu(x) x nn.max_pool(x, window_shape(2, 2), strides(2, 2)) x nn.Conv(features64, kernel_size(3, 3), paddingSAME)(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(features128)(x) x nn.relu(x) x nn.Dense(features10)(x) return x這里使用 Flax 來定義模型。Flax 是 Google 官方維護(hù)的 JAX 神經(jīng)網(wǎng)絡(luò)庫和 JAX 一起使用時(shí)體驗(yàn)最自然。如果你之前用過 PyTorch 的nn.ModuleFlax 的nn.Module風(fēng)格會(huì)比較接近。創(chuàng)建train.py這是訓(xùn)練循環(huán)的核心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_classes10) 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, trainingTrue) return cross_entropy_loss(logits, labels) loss, grads jax.value_and_grad(loss_fn)(state.params) state state.apply_gradients(gradsgrads) return state, loss jax.jit def eval_step(state, batch): images, labels batch logits state.apply_fn({params: state.params}, images, trainingFalse) 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_fnmodel.apply, paramsparams, txtx) def main(): rng jax.random.PRNGKey(0) state create_train_state(rng, learning_rate1e-3) train_ds, test_ds load_mnist(batch_size128) # 將 TensorFlow Dataset 轉(zhuǎn)換為 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(fepoch {epoch} step {step} loss {loss:.4f}) # 每個(gè) epoch 結(jié)束評(píng)估一次 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(fepoch {epoch} eval loss {total_loss / num_batches:.4f} facc {total_acc / num_batches:.4f}) if __name__ __main__: main()這段代碼有幾個(gè)關(guān)鍵點(diǎn)使用jax.jit裝飾訓(xùn)練和評(píng)估函數(shù)讓 XLA 將整個(gè)計(jì)算過程編譯成融合算子。每個(gè) batch 通過.numpy()從 TensorFlow Dataset 轉(zhuǎn)換為 NumPy 數(shù)組供 JAX 消費(fèi)。訓(xùn)練狀態(tài)由TrainState統(tǒng)一管理包含模型參數(shù)和優(yōu)化器狀態(tài)。打印網(wǎng)絡(luò)在 MNIST 上的 loss 和 acc用來驗(yàn)證模型真實(shí)地訓(xùn)練起來了。4.4 真機(jī)運(yùn)行與驗(yàn)證在 Colab 或 TPU VM 上執(zhí)行python3 train.py正常輸出會(huì)類似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 穩(wěn)步上升說明 JAX XLA TPU 這條鏈路已經(jīng)跑通了。接下來就可以把這里的SimpleCNN替換成真實(shí)模型把load_mnist替換成你的真實(shí)數(shù)據(jù)集。4.5 關(guān)于多卡 TPU 的擴(kuò)展思路上面的代碼是單進(jìn)程、單 TPU 核心訓(xùn)練的寫法。如果你創(chuàng)建的是多核心 TPUJAX 會(huì)自動(dòng)識(shí)別多個(gè)設(shè)備。想要利用多個(gè)核心并行訓(xùn)練最簡(jiǎn)單的方式是使用jax.device_put和pmap對(duì) batch 做數(shù)據(jù)并行切分。from jax import pmap # 將狀態(tài)復(fù)制到所有設(shè)備 state jax.device_put_replicated(state, jax.devices()) # 多設(shè)備并行訓(xùn)練步驟 pmap def train_step_multi(state, batch): return train_step(state, batch) # 切分 batch 到多個(gè)設(shè)備 images images.reshape((num_devices, -1) images.shape[1:]) labels labels.reshape((num_devices, -1))這只是一個(gè)非常簡(jiǎn)化的pmap示例真實(shí)多卡訓(xùn)練還要考慮 batch 切分策略、梯度累積、AllReduce 等細(xì)節(jié)。建議先把單 core 代碼跑通再逐步過渡到pmap或shard_map。5. TensorFlow 方式用 TPUStrategy 做分布式訓(xùn)練雖然 JAX 在 TPU 上體驗(yàn)越來越主流但生產(chǎn)環(huán)境中大量存量代碼還是 TensorFlow 的。如果你不想重寫模型可以直接使用 TensorFlow 的TPUStrategy把現(xiàn)有模型遷移到 TPU。5.1 TPUStrategy 工作原理TPUStrategy是 TensorFlow 的分布式策略之一。它會(huì)自動(dòng)完成幾件事把模型變量復(fù)制到每個(gè) TPU 核心。把全局 batch 切分成每個(gè)核心處理一個(gè)子 batch。優(yōu)化器在反向傳播后做跨核心梯度 AllReduce。對(duì)訓(xùn)練循環(huán)內(nèi)部使用tf.function編譯配合 XLA 加速。它的關(guān)鍵思路是“數(shù)據(jù)并行 同步更新”。用戶只需要寫一份單機(jī)單卡代碼策略層負(fù)責(zé)把計(jì)算擴(kuò)展到多個(gè) TPU 核心上。5.2 TPUStrategy 訓(xùn)練示例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 作用域內(nèi)構(gòu)建模型 with strategy.scope(): model tf.keras.Sequential([ tf.keras.layers.Conv2D(32, (3, 3), activationrelu, input_shape(28, 28, 1)), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.Conv2D(64, (3, 3), activationrelu), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.Flatten(), tf.keras.layers.Dense(128, activationrelu), tf.keras.layers.Dense(10), ]) model.compile( optimizertf.keras.optimizers.Adam(1e-3), losstf.keras.losses.SparseCategoricalCrossentropy(from_logitsTrue), metrics[accuracy], ) # 3. 準(zhǔn)備數(shù)據(jù) (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. 訓(xùn)練 model.fit( x_train, y_train, batch_size128, epochs3, validation_data(x_test, y_test), )TPUStrategy的遷移成本很低尤其適合原本就是 Keras 寫的模型。如果你的模型里含有自定義tf.Variable更新邏輯就需要確認(rèn)變量創(chuàng)建是否都發(fā)生在strategy.scope()內(nèi)。5.3 TFRecord 數(shù)據(jù)管道與 TPU 的配合真實(shí)訓(xùn)練中數(shù)據(jù)通常在 GCS 上推薦使用 TFRecord 格式來減少小文件數(shù)量提高 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_trainingTrue): dataset tf.data.Dataset.list_files(file_pattern) dataset dataset.interleave( tf.data.TFRecordDataset, cycle_length8, num_parallel_callstf.data.AUTOTUNE, ) dataset dataset.map(decode_fn, num_parallel_callstf.data.AUTOTUNE) if is_training: dataset dataset.shuffle(10000).repeat() dataset dataset.batch(batch_size, drop_remainderTrue) dataset dataset.prefetch(tf.data.AUTOTUNE) return dataset數(shù)據(jù)管道本身是 CPU 上的工作但它的吞吐能力直接決定了 TPU 是否被喂飽。經(jīng)驗(yàn)法則是數(shù)據(jù)管道的吞吐至少要達(dá)到模型訓(xùn)練速度的 2 倍否則訓(xùn)練過程中會(huì)不斷出現(xiàn)“等待數(shù)據(jù)”的空窗期。6. 常見問題與排查思路6.1 TPU 初始化失敗或設(shè)備不可見問題現(xiàn)象RuntimeError: Unable to initialize the TPU system.可能原因當(dāng)前運(yùn)行時(shí)不是 TPU 環(huán)境。JAX/TensorFlow 版本與 TPU 運(yùn)行時(shí)版本不匹配。運(yùn)行時(shí)使用舊內(nèi)核沒有加載必要的驅(qū)動(dòng)。排查步驟先檢查設(shè)備文件是否存在在 TPU VM 上執(zhí)行l(wèi)s /dev/accel*。再執(zhí)行jax.devices()或tf.config.list_logical_devices(TPU)看框架層是否識(shí)別到設(shè)備。確認(rèn)虛擬環(huán)境中的包版本與當(dāng)前 TPU 版本匹配。重啟運(yùn)行時(shí)尤其是在切換 TPU 類型之后。6.2 XLA 編譯報(bào)錯(cuò) “unsupported operations”問題現(xiàn)象Detected unsupported operations when trying to compile graph可能原因模型里使用了 XLA 不支持的算子。自定義 Layer 或自定義 JAX 函數(shù)里調(diào)用了無法被轉(zhuǎn)換成 HLO 的 Python 控制流。解決思路逐步裁剪模型定位到具體是哪個(gè)操作不能被編譯。檢查是否是動(dòng)態(tài)形狀問題TPU 上盡量避免 runtime shape 變化。對(duì)自定義算子優(yōu)先考慮能否用現(xiàn)有的 JAX 原生操作重寫。6.3 關(guān)于 16KB page size 編譯包報(bào)錯(cuò)問題現(xiàn)象在安裝或運(yùn)行 TPU SDK 相關(guān)包時(shí)出現(xiàn)類似 “an error occurred while preparing sdk package 16 kb page size” 的報(bào)錯(cuò)。背景原因TPU 運(yùn)行環(huán)境的虛擬地址頁大小是 16KB而常見 x86 Linux 是 4KB 頁。本地下載的預(yù)編譯 wheel 如果是在 4KB 頁環(huán)境下構(gòu)建的運(yùn)行時(shí)會(huì)與 TPU 內(nèi)核模塊或運(yùn)行時(shí)庫不兼容。解決思路不要直接從普通 PyPI 鏡像安裝 TPU 相關(guān)包優(yōu)先使用官方發(fā)布渠道。在 TPU VM 內(nèi)重新安裝匹配版本而不是把本地環(huán)境復(fù)制過去。查看完整日志確認(rèn)是下載失敗還是安裝后運(yùn)行時(shí)崩潰。6.4 OOM 與 batch size 調(diào)整問題現(xiàn)象訓(xùn)練時(shí)出現(xiàn)Resource exhausted錯(cuò)誤??赡茉騜atch size 過大超出了 TPU 單核的 HBM 容量。模型參數(shù)或中間激活值過大。數(shù)據(jù)管道prefetch使用內(nèi)存過多。解決思路先嘗試減小 batch size確認(rèn)模型在單卡上可以跑通。合理設(shè)置prefetch和num_parallel_calls避免 CPU 端內(nèi)存溢出。檢查 XLA 是否啟用了內(nèi)存優(yōu)化選項(xiàng)但不要盲目相信默認(rèn)配置。6.5 排查清單匯總問題現(xiàn)象常見原因解決思路設(shè)備不可見運(yùn)行時(shí)環(huán)境錯(cuò)誤 / 版本不匹配檢查/dev/accel*重啟運(yùn)行時(shí)核對(duì)版本XLA 編譯失敗不支持的算子 / 動(dòng)態(tài) shape簡(jiǎn)化模型逐步定位避免動(dòng)態(tài)形狀16KB page size 報(bào)錯(cuò)wheel 與 TPU 頁大小不匹配使用官方發(fā)布渠道安裝匹配包OOMbatch size 過大 / 中間激活過大調(diào)整 batch size優(yōu)化數(shù)據(jù)管道內(nèi)存訓(xùn)練速度上不去數(shù)據(jù)管道存在瓶頸使用 TFRecord、interleave、prefetch 優(yōu)化輸出結(jié)果有 NaN學(xué)習(xí)率過高 / 精度配置不當(dāng)降低學(xué)習(xí)率使用混合精度時(shí)檢查損失縮放7. 最佳實(shí)踐與工程建議7.1 數(shù)據(jù)管道要提前壓測(cè)TPU 非常“挑食”它的運(yùn)算速度快到如果數(shù)據(jù)管道沒跟上訓(xùn)練就會(huì)空等。建議在正式訓(xùn)練前先單獨(dú)測(cè)試數(shù)據(jù)管道的吞吐能力。import time train_iter iter(train_ds) start time.time() for i in range(100): batch next(train_iter) end time.time() print(f100 batches take {end - start:.2f}s)如果數(shù)據(jù)讀取時(shí)間占比過高優(yōu)先使用 TFRecord 格式、interleave并行讀取、prefetch預(yù)加載等手段。7.2 使用混合精度時(shí)需要驗(yàn)證損失縮放TPU 的 bfloat16 支持是它的強(qiáng)項(xiàng)之一。用jax時(shí)可以很自然地把部分參數(shù)轉(zhuǎn)換成bfloat16或使用混合精度訓(xùn)練。但在混合精度下梯度很小的時(shí)候有可能在低精度下溢出導(dǎo)致 NaN。建議開啟損失縮放Loss Scaling機(jī)制并周期性檢查梯度統(tǒng)計(jì)而不是在出現(xiàn) NaN 后才去排查。7.3 模型算子優(yōu)先考慮 JAX/TensorFlow 原生實(shí)現(xiàn)遇到自定義算子時(shí)優(yōu)先檢查原生庫里有沒有替代實(shí)現(xiàn)。XLA 對(duì)原生算子的融合優(yōu)化做得很成熟但自定義算子往往無法被融合會(huì)打亂 XLA 的優(yōu)化策略。如果一定要使用自定義算子至少把“無法編譯”的算子獨(dú)立出來避免拖累整體性能。7.4 成本與資源管理TPU 資源通常是按時(shí)計(jì)費(fèi)的長(zhǎng)期任務(wù)建議配合 checkpoint 和自動(dòng)重啟機(jī)制。這里給出幾個(gè)實(shí)用性建議訓(xùn)練腳本要有穩(wěn)定的 checkpoint 保存與恢復(fù)邏輯。使用搶占式資源時(shí)代碼要能安全處理進(jìn)程中途被回收的情況。創(chuàng)建 TPU 實(shí)例后盡快驗(yàn)證訓(xùn)練腳本避免空轉(zhuǎn)計(jì)費(fèi)。配置監(jiān)控報(bào)警對(duì)異常掉線、訓(xùn)練停滯、指標(biāo)不回傳這些情況做告警。7.5 把訓(xùn)練過程日志化TPU 上的日志查看比 GPU 環(huán)境要麻煩一些尤其是多 worker 場(chǎng)景下。推薦在訓(xùn)練腳本中顯式記錄關(guān)鍵指標(biāo)比如每個(gè) epoch 的 loss、acc、每個(gè) step 的平均耗時(shí)、數(shù)據(jù)加載耗時(shí)等并把日志輸出到統(tǒng)一的日志平臺(tái)方便后續(xù)歸因。import logging logging.basicConfig( levellogging.INFO, format%(asctime)s %(levelname)s %(message)s, ) logging.info(start training, devices%s, jax.devices())這一步可能看起來不起眼但當(dāng)你遇到 TPU 訓(xùn)練中途卡死、重啟、性能退化問題時(shí)這些日志是定位問題的唯一線索。8. 總結(jié)Google TPU 軟件棧本質(zhì)上是一條從深度學(xué)習(xí)框架到專用硬件的編譯和運(yùn)行鏈路JAX、TensorFlow、XLA 和 libtpu 各司其職。對(duì)開發(fā)者而言遷移到 TPU 時(shí)最需要調(diào)整的往往不是模型結(jié)構(gòu)本身而是對(duì)計(jì)算圖編譯、數(shù)據(jù)切分、內(nèi)存規(guī)劃這些底層層面的理解。如果你準(zhǔn)備開始嘗試建議從 Colab 免費(fèi) TPU 環(huán)境入手把 MNIST 或自己的小型模型跑通再逐步擴(kuò)展到真實(shí)業(yè)務(wù)模型。遇到 16KB page size、XLA 編譯失敗、TPU 設(shè)備不可見這類問題先順著軟件棧的層級(jí)逐層排查不要一開始就懷疑硬件。只要把 JAX 或 TensorFlow 的 TPU 鏈路跑通一遍后續(xù)在 Cloud TPU 上做規(guī)?;?xùn)練會(huì)順暢得多。