)
這次我們來看 Google TPU 軟件棧。注意這不是一篇 TPU 芯片參數(shù)測評而是把你在 TPU 上實際會用到的那一套工具鏈講清楚JAX、TensorFlow、PyTorch/XLA、XLA 編譯器、PJRT 運行時以及它們怎么配合把一個模型的訓練任務從 GPU 平移到 TPU再推上生產(chǎn)環(huán)境。為什么值得關注 TPU 軟件棧因為 TPU 和 GPU 的編程模型差異并不在框架層而在編譯器層和運行時層。GPU 生態(tài)由 CUDA 統(tǒng)治TPU 生態(tài)則圍繞 XLA 和 PJRT 構建。JAX、TensorFlow、PyTorch 都能在 TPU 上跑但前提是正確安裝軟件棧、正確編寫模型代碼、正確布局數(shù)據(jù)并行和模型并行。這些內(nèi)容不是簡單把device(cuda)改成device(tpu)就能解決的。本文會帶你過一遍TPU 軟件棧的組成、在 Google Cloud 上開通環(huán)境、在 Colab/Kaggle 免費入口快速上手、用 JAX 訓練一個小模型、把 PyTorch 模型遷移到 TPU、批量提交訓練任務、觀察資源占用、排查常見問題。適合正在做 AI 基礎設施選型、想把模型訓練規(guī)?;?、或者計劃從 GPU 遷移到 TPU 的工程師。如果只是想在低成本算力上跑個 DemoColab 和 Kaggle 的免費 TPU 入口也足夠你體驗完整流程。1. Google TPU 軟件棧核心能力速覽能力項說明項目類型AI 計算加速硬件 配套軟件棧主要組件XLA 編譯器、JAX、TensorFlow、PyTorch/XLA、PJRT、libTPU主要功能大規(guī)模模型訓練、分布式訓練、推理加速支持數(shù)據(jù)并行和模型并行編程接口JAX API、TensorFlow API、PyTorch/XLA訪問方式Google Cloud 云服務Colab / Kaggle 提供實驗性質免費入口硬件門檻TPU 在云端訪問本機只需能跑框架客戶端本地仿真不需要 TPU 硬件顯存 / HBM不同代際 TPU 芯片配備不同容量 HBM以官方文檔和實際規(guī)格為準啟動方式創(chuàng)建 Cloud TPU VM 后 SSH 進入或通過 Vertex AI、Notebook 提交任務是否支持批量支持可同時創(chuàng)建 TPU 切片、批量提交訓練任務是否支持 API不直接提供推理 API通常配合 Vertex AI 或自建服務暴露模型端點適用場景大規(guī)模訓練、超參搜索、多租戶模型實驗、生產(chǎn)推理這里需要先區(qū)分一個概念TPU 不是像顯卡一樣插在自己機器上的加速卡而是 Google Cloud 上的托管硬件。你使用 TPU 的方式是先創(chuàng)建 TPU VM再在 VM 里裝框架、寫代碼、跑任務。所以“本地部署”在 TPU 場景里更多指“本機作為控制端 云端算力”。2. 為什么 AI 規(guī)?;涞匦枰?TPU 軟件棧2.1 AI 時代對芯片的新需求搜索熱詞里經(jīng)常能看到“AI 時代對于芯片的特殊需求”這類討論。芯片行業(yè)過去幾十年主要圍繞 CPU 的通用計算能力迭代但 AI 負載的核心是矩陣乘法、卷積、attention 這類密集張量運算。CPU 擅長復雜分支和邏輯控制但矩陣運算效率并不理想GPU 通過大規(guī)模并行線程解決了吞吐問題但驅動復雜、功耗高、供應緊張NPU 更多出現(xiàn)在移動端和邊緣側。TPU 則是 Google 為 AI 計算設計的專用加速器核心目標很直接把矩陣乘法這類操作以最有效的方式算完。規(guī)?;涞貢r問題往往不是“芯片不夠快”而是“軟件棧能不能把硬件發(fā)揮出來”。很多團隊買來算力后發(fā)現(xiàn)模型遷移、分布式并行、性能調(diào)優(yōu)才是真正耗時間的部分。TPU 軟件棧存在的意義就是把硬件能力通過一套標準化接口暴露給上層框架降低遷移成本。2.2 TPU 軟件棧解決三件事第一編譯層。XLA 編譯器把模型計算圖轉換成 TPU 可執(zhí)行的低層指令。你在 JAX 里寫的jit、在 PyTorch/XLA 里的mark_step或torch.compile底層都會走 XLA。沒有 XLA任何框架都無法在 TPU 上高效執(zhí)行。第二運行時層。PJRTPortable JAX Runtime統(tǒng)一了框架側到設備側的調(diào)用路徑。它負責設備發(fā)現(xiàn)、程序加載、執(zhí)行、內(nèi)存管理和集合通信。JAX、TensorFlow、PyTorch/XLA 共用同一套 PJRT 運行時這是 TPU 能跨框架工作的關鍵設計。第三框架層。JAX 是 Google 為科研和訓練設計的數(shù)組計算庫支持自動微分和 GPU/TPU 無縫切換TensorFlow 是老牌訓練框架PyTorch/XLA 則讓 PyTorch 生態(tài)也能在 TPU 上運行。也就是說團隊現(xiàn)有的 PyTorch 代碼不用重寫而是切換設備字符串加少量適配。3. TPU 軟件棧的組成與工作原理3.1 TPU 硬件的基本劃分先了解幾個名詞后面講軟件棧時會反復出現(xiàn)。TPU 有幾種常見形態(tài)v2/v3 是早期 Pod 形態(tài)適合小規(guī)模實驗v4 是分代產(chǎn)品算力更強v5e 主打性價比適合中小規(guī)模訓練和推理v5p 面向更大規(guī)模的訓練負載Trilliumv6是新一代 TPU在整數(shù)性能、稀疏性和配置靈活性上做了升級。具體代際參數(shù)不用死記重點是每一塊 TPU 設備有獨立的計算單元和 HBM 容量多塊卡通過高速互連組成 TPU slice切片軟件層面看到的是一個分布式設備集群而不是多塊獨立顯卡。這也意味著TPU 軟件棧里的分布式通信層非常重要。你寫代碼時可能只看到一個jax.devices()返回多設備列表但底層已經(jīng)幫你處理了設備間通信拓撲。3.2 XLA 編譯器XLA 的全稱是 Accelerated Linear Algebra。它的工作分幾步接收來自 JAX / TensorFlow / PyTorch 的高層計算圖通常是 HLO 表示。做算子融合。把多個小算子合并成大的融合算子減少內(nèi)存訪問次數(shù)。做布局優(yōu)化。為 TPU 的內(nèi)存排布選擇最優(yōu)的張量布局。生成低層可執(zhí)行代碼交給 TPU 運行。從用戶視角看最重要的是你寫的模型如果不經(jīng)過 XLATPU 根本跑不了。JAX 里用jax.jit編譯函數(shù)TensorFlow 2.x 的 Graph 模式默認走 XLAPyTorch/XLA 里torch.compile或mark_step都會觸發(fā)編譯。所以當模型運行速度很慢時首先要確認編譯是否真的發(fā)生。3.3 PJRT 運行時PJRT 統(tǒng)一了框架側到設備側的調(diào)用路徑。以前 JAX 和 TensorFlow 各自維護一套運行時后來 OpenXLA 項目把這些邏輯統(tǒng)一到 PJRT 里。PJRT 負責設備發(fā)現(xiàn)列出當前可用的 TPU 設備負責程序加載把 XLA 編譯結果發(fā)到 TPU負責執(zhí)行同步或異步運行計算負責內(nèi)存管理分配和釋放設備內(nèi)存還提供 all-reduce、all-gather、reduce-scatter 等集合通信原語。由于 PJRT 存在PyTorch 社區(qū)和 JAX 社區(qū)可以共用同一個 TPU 運行時。這也是為什么 PyTorch 模型切到 TPU 時只需要安裝 torch-xla 并設置環(huán)境變量而不是重新實現(xiàn)一套設備棧。3.4 數(shù)據(jù)并行與模型并行TPU 軟件棧天然支持兩種并行方式數(shù)據(jù)并行多個 TPU 設備持有同一份模型參數(shù)每個設備處理不同的 batch定期 all-reduce 梯度。模型并行 / 張量并行模型太大放不進單設備 HBM 時把參數(shù)切分到多臺設備上。JAX 里通過jax.sharding描述數(shù)據(jù)如何分布到設備PyTorch 里使用 FSDP 或 DDP 的 XLA 適配TensorFlow 里使用 TF Distribution Strategy。軟件棧會把這些策略翻譯成 XLA 編譯行為再通過 PJRT 的集合通信完成梯度和權重同步。理解這一點你才能解釋同一個模型在單卡和切片上為什么性能差距很大。4. 適用場景與使用邊界4.1 適合誰正在做大規(guī)模訓練的團隊是最直接的受益者。單機 GPU 顯存不足時TPU slice 能提供更大的單任務算力而且不需要自己搭 InfiniBand 網(wǎng)絡。已經(jīng)用 JAX 或 TensorFlow 的團隊遷移成本相對低因為這兩套框架與 TPU 軟件棧本身就是同時演進的。想用低成本算力做實驗的開發(fā)者也可以關注 Colab 和 Kaggle 的免費 TPU 入口雖然不適合生產(chǎn)但足夠驗證模型邏輯和跑通訓練鏈路。想對比 GPU 和 TPU 訓練成本的團隊建議先拿典型 workload 做基準測試。TPU 在矩陣密集型任務里效率很高但在小 batch、強數(shù)據(jù)增強、頻繁動態(tài) shape 的場景下不一定比高端 GPU 有優(yōu)勢。是否選擇 TPU取決于你的模型結構、數(shù)據(jù)管道和預算結構。4.2 不適合誰如果只需要一張小顯存顯卡跑個小 Demo本地 GPU 可能更省事沒必要引入云端 TPU 的編排開銷。如果項目強依賴 CUDA 生態(tài)比如某些只支持 CUDA 的第三方算子遷移到 TPU 需要重新驗證和實現(xiàn)。如果模型很小、延遲要求極高TPU 的調(diào)度和網(wǎng)絡開銷也不一定比 GPU 好。4.3 合規(guī)與安全邊界使用 Google TPU 必須通過 Google Cloud 正規(guī)渠道遵守當?shù)胤煞ㄒ?guī)和 Google Cloud 服務條款。以下幾點需要特別注意不要共享云賬號、密鑰或未經(jīng)授權的 TPU 資源。訓練數(shù)據(jù)、模型權重如果是企業(yè)資產(chǎn)或涉及隱私要確認數(shù)據(jù)駐留和數(shù)據(jù)治理策略。生成的模型和輸出內(nèi)容不能用于侵權、造假、欺騙等非法用途。如果準備用 TPU 跑換臉、聲音克隆等內(nèi)容必須先確認素材授權和肖像、聲音權利否則不要碰。Colab 和 Kaggle 的免費 TPU 入口有配額限制只能用于學習與實驗。5. TPU 軟件棧環(huán)境準備與前置條件5.1 云上環(huán)境準備正式使用 Cloud TPU需要在 Google Cloud 上準備好項目、API 和配額。你需要一個 Google Cloud 賬號創(chuàng)建一個項目啟用 Cloud TPU API確認目標區(qū)域有 TPU 資源配額。如果區(qū)域沒有可用配額需要申請或換區(qū)域。本機不需要高配置重點是有權限訪問 Google Cloud 服務并按你所在地區(qū)的法律法規(guī)和公司網(wǎng)絡策略來操作這里不展開。5.2 本地仿真與免費入口生產(chǎn)前想快速體驗 TPU 軟件棧有兩個免費入口Colab新建筆記本后在運行時類型里選擇 TPU就能得到一個 TPU 運行時。適合跑 JAX 和 TensorFlow。KaggleKaggle Notebook 也提供 TPU 加速器選項配合公開數(shù)據(jù)集做實驗很方便。這類免費環(huán)境不是 Cloud TPU 的生產(chǎn)形態(tài)但可以完成軟件棧體驗、代碼驗證和模型調(diào)試。正式訓練建議用 Cloud TPU配額和性能更有保證。5.3 磁盤空間和依賴TPU VM 默認有一個系統(tǒng)盤模型和數(shù)據(jù)集最好放在持久磁盤或 Cloud Storage避免 VM 重建時丟失。需要安裝的依賴包括 Python 3.9、JAX TPU 版、PyTorch、torch-xla、TensorFlow。版本和 Cloud TPU VM 官方鏡像保持一致。一個通用做法是在 TPU VM 里創(chuàng)建 Python 虛擬環(huán)境隔離項目依賴避免把系統(tǒng) Python 搞亂。6. TPU 軟件棧安裝部署與啟動方式6.1 創(chuàng)建 TPU VM用 gcloud 創(chuàng)建 Cloud TPU 的通用命令模板如下具體參數(shù)需要按你的項目、區(qū)域、TPU 型號替換。不同區(qū)域可用的 TPU 型號和版本鏡像可能不同以實際配額和官方文檔為準。# 先配置項目 gcloud config set project YOUR_PROJECT_ID # 創(chuàng)建 TPU VM這里以 v5e-4 為例實際型號以可用配額為準 gcloud compute tpus tpu-vm create tpu-demo \ --zoneus-central1-b \ --accelerator-typev5e-4 \ --versiontpu-vm-v4-base \ --projectYOUR_PROJECT_ID創(chuàng)建完成后用 SSH 進入 TPU VMgcloud compute tpus tpu-vm ssh tpu-demo --zoneus-central1-b進入之后可以先確認設備是否正常ls /dev/accel* python3 -c import jax; print(jax.devices())如果看到TpuDevice列表說明 PJRT 運行時已經(jīng)能識別 TPU。如果沒有說明 JAX 安裝版本或 PJRT 環(huán)境變量有問題。6.2 啟動一個 JAX 驗證腳本在 TPU VM 上最簡單的驗證方式是創(chuàng)建虛擬環(huán)境并安裝 JAXpython3 -m venv venv source venv/bin/activate pip install --upgrade pip # JAX 官方 TPU 版安裝方式版本以官方發(fā)布為準 pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html安裝完成后運行最小驗證腳本import jax import jax.numpy as jnp print(device count:, jax.device_count()) print(devices:, jax.devices()) # 在 TPU 上做矩陣乘法 a jnp.ones((4096, 4096), dtypejnp.bfloat16) b jnp.ones((4096, 4096), dtypejnp.bfloat16) c jnp.matmul(a, b) # 觸發(fā)計算并拉回 CPU result jax.device_get(c) print(done, result shape:, result.shape, first value:, result[0, 0])如果腳本順利輸出設備列表和矩陣乘法結果說明 JAX、XLA、PJRT 到 TPU 的整條鏈路已經(jīng)打通。接下來就可以在這個基礎上替換數(shù)據(jù)集、模型結構逐步接近真實訓練任務。6.3 啟動一個 PyTorch/XLA 腳本PyTorch 遷移到 TPU 的核心是安裝 torch-xla并設置 XLA 環(huán)境變量pip install torch torch-xla export PJRT_DEVICETPU export XLA_USE_BF161運行 PyTorch 的 TPU 驗證腳本import torch import torch_xla import torch_xla.core.xla_model as xm device xm.xla_device() print(device:, device) x torch.randn(8, 8, devicedevice) y torch.randn(8, 8, devicedevice) z x y # XLA 懶執(zhí)行必須觸發(fā)一次 step 才能真正完成計算 xm.mark_step() print(result:, z)這里涉及一個核心概念XLA 默認懶執(zhí)行xm.mark_step()或xm.xla_step()觸發(fā)一次設備執(zhí)行。在訓練循環(huán)里每步都要調(diào)用mark_step否則數(shù)據(jù)不會真正送到 TPU。很多 PyTorch 用戶剛遷移過來忘記這一步導致訓練循環(huán)不執(zhí)行或者 loss 原地不動屬于最常見的坑。7. TPU 軟件棧功能測試與效果驗證7.1 JAX 訓練一個小模型用簡單的 MLP 模型完整跑一次訓練既能驗證軟件棧又能驗證數(shù)據(jù)搬運、梯度和優(yōu)化器是否正常工作。import jax import jax.numpy as jnp from jax import random # 隨機生成小分類數(shù)據(jù)集 key random.PRNGKey(0) X random.normal(key, (1024, 784)) y random.randint(key, (1024,), 0, 10) # 初始化參數(shù) def init_params(rng): k1, k2 random.split(rng) w1 random.normal(k1, (784, 128)) * 0.1 w2 random.normal(k2, (128, 10)) * 0.1 return w1, w2 # 前向 def forward(params, x): w1, w2 params h jnp.tanh(x w1) return h w2 # 損失 def loss_fn(params, x, y): logits forward(params, x) one_hot jax.nn.one_hot(y, num_classes10) return jnp.mean(jnp.sum(-one_hot * jax.nn.log_softmax(logits), axis-1)) # 訓練步 jax.jit def train_step(params, x, y, lr0.01): grads jax.grad(loss_fn)(params, x, y) return [(p - lr * g) for p, g in zip(params, grads)] params init_params(key) for epoch in range(5): params train_step(params, X, y) loss loss_fn(params, X, y) print(fepoch {epoch}, loss: {float(loss):.4f})jax.jit會觸發(fā) XLA 編譯并緩存計算圖訓練循環(huán)里每一輪實際跑的是編譯后的 TPU 程序。如果這個腳本完整跑通說明 JAX 梯度計算、jit 編譯、張量調(diào)度都沒有問題。第一次運行會看到編譯耗時較長第二次開始會明顯變快這是正常現(xiàn)象。7.2 PyTorch/XLA 訓練驗證PyTorch 模型遷移到 TPU 的關鍵改動點有三個設備從cuda改成xla_device()訓練循環(huán)里顯式調(diào)用xm.mark_step()梯度更新使用xm.optimizer_step替代普通optimizer.step()。import torch import torch.nn as nn import torch.optim as optim import torch_xla import torch_xla.core.xla_model as xm device xm.xla_device() model nn.Sequential( nn.Linear(784, 128), nn.Tanh(), nn.Linear(128, 10), ).to(device) optimizer optim.SGD(model.parameters(), lr0.01) loss_fn nn.CrossEntropyLoss() x torch.randn(128, 784, devicedevice) y torch.randint(0, 10, (128,), devicedevice) for step in range(50): optimizer.zero_grad