
做深度學習第一步不是搭模型而是選框架。PyTorch、TensorFlow、JAX 這三套 API 設計思路差異非常大代碼遷移成本高選錯了后面寫訓練循環、做部署都要返工。這篇文章以“7 個框架全景”為背景重點把三個主力框架的核心 API 拆開對比環境安裝、張量操作、自動微分、模型構建、訓練循環、數據加載、部署導出、資源占用和問題排查全部用最小可運行示例驗證。如果你正打算從 TensorFlow 遷到 PyTorch或者想試 JAX 的函數式變換但一直沒上手這篇可以直接收藏。先給結論沒有“最強框架”只有“最匹配場景的 API”。PyTorch 適合研究和快速迭代TensorFlow 適合生產系統和端側部署JAX 適合需要自動微分、自動向量化、自動并行化的高性能數值計算。下面用一套“從環境到推理”的最小測試流程把三個框架的差異點逐個講清楚。1. 7 個深度學習框架核心能力速覽框架來源/社區核心 API 風格主要定位PyTorchMeta 發起現由 PyTorch Foundation 管理動態圖nn.Module autograd科研、快速原型、工業推理TensorFlowGoogle動態圖 Keras 高層 API生產系統、端側部署、大規模分布式JAXGoogle函數式變換jnp grad/jit/vmap/pmap高性能數值計算、科研算法復現KerasGoogle / 社區高層 API支持多后端快速搭建、教學、遷移到不同后端PaddlePaddle百度動態圖為主動靜統一工業應用與科研中文生態MindSpore華為動靜統一自動并行昇騰硬件生態、企業 AIMXNetApacheGluon 動態圖接口歷史使用廣泛當前社區活躍度明顯下降這 7 個框架里PyTorch、TensorFlow、JAX 的開源社區最活躍也是國內外論文復現和工程落地的絕對主流。Keras 現在更像一個“前端接口層”可以跑在 TensorFlow、JAX 甚至 PyTorch 后端上。PaddlePaddle、MindSpore 在特定硬件和中文工業場景里有很強支持。MXNet 今天主要用于維護老項目新項目不建議再選。2. 框架選型什么時候用哪個選框架不要看熱度要看你要做什么。PyTorch 最值得選的理由是“調試直接”。模型就是一個普通 Python 對象前向傳播是一段普通 Python 代碼print、breakpoint、pdb 都能直接用。論文復現、快速驗證新想法、做多輪實驗PyTorch 的效率最高。PyTorch 的模型結構定義和動態控制流非常自然RNN、Transformer、擴散模型這類結構寫起來都不費勁。TensorFlow 的強項在“生產鏈路完整”。從 TF Serving、TFLite、TensorFlow.js 到 TFX 流水線訓練到部署的工程件齊全。Keras 高層接口讓模型搭建非常快適合標準化團隊協作和產品化落地。如果團隊已經有完整的 Kubernetes 和模型服務基礎設施TensorFlow 仍然是可靠選擇。JAX 要接受的是一套完全不同的思維沒有“模型對象”一切是純函數加參數數組沒有model.fit()訓練循環要自己寫換來的好處是grad、jit、vmap、pmap這種組合式變換在強化學習、分子動力學、貝葉斯建模、大模型分布式并行等場景里效率極高。Google DeepMind 的很多研究項目和開源庫都基于 JAX。使用邊界也要說清楚訓練數據必須來源合法模型權重要看開源許可證涉及人臉、聲音、隱私數據時要確認授權部署 API 服務要限制訪問范圍任何情況下都不要用框架去繞過安全限制或做侵權內容生成。3. 環境準備與安裝三套框架都要求先確認 Python、CUDA、cuDNN 版本匹配。裝不上 GPU 版通常是 CUDA 版本不一致或者驅動太舊排查順序固定為驅動 - CUDA - Python - 框架。通用檢查命令python --version nvidia-smi nvcc --version驅動只要滿足 CUDA 版本即可。注意nvidia-smi顯示的是驅動支持的 CUDA 版本不一定是本機安裝的 CUDA toolkit 版本兩個概念不要混淆。PyTorch 安裝推薦用官方命令生成器# 不要直接復制去 pytorch.org 根據系統、CUDA 版本生成命令 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121TensorFlow 2.18 的 Linux GPU 安裝推薦元包方式# Linux GPU 版CUDA 相關依賴由元包統一管理 pip install tensorflow[and-cuda] # 僅 CPU 版 pip install tensorflow-cpuJAX 的 GPU 安裝要區分 CUDA 版本# CPU 版 pip install -U jax # CUDA 12 版 pip install -U jax[cuda12] # CUDA 11 版 pip install -U jax[cuda11]安裝完不要急著寫模型先做硬件驗證# PyTorch import torch print(PyTorch, torch.__version__) print(CUDA available:, torch.cuda.is_available()) print(GPU:, torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU) # TensorFlow import tensorflow as tf print(TensorFlow, tf.__version__) print(GPU:, tf.config.list_physical_devices(GPU)) # JAX import jax print(JAX, jax.__version__) print(Devices:, jax.devices())JAX 如果輸出CpuDevice說明 GPU 驅動或 CUDA 庫沒配對。TensorFlow 2.18 如果只顯示 CPU重點檢查tensorflow[and-cuda]是否安裝成功而不是只裝了tensorflow基礎包。4. 核心張量 API 對比三者的核心張量類型分別是torch.Tensor、tf.Tensor、jax.ArrayJAX 統一用jnp.array創建。API 設計差異在創建、轉換、設備管理上非常明顯。import torch import tensorflow as tf import jax.numpy as jnp # PyTorch pt torch.tensor([1.0, 2.0, 3.0]) print(pt.dtype, pt.device) # TensorFlow tf_t tf.constant([1.0, 2.0, 3.0]) print(tf_t.dtype, tf_t.device) # JAX jx jnp.array([1.0, 2.0, 3.0]) print(jx.dtype, jx.device())三者的設計差異PyTorch 的張量是“可變的”。你可以原地改數據、移動設備但自動微分需要參與求導的張量顯式設置requires_gradTrue默認關閉。TensorFlow 的張量默認“不可變”但tf.Variable是可變對象。Keras 模型里的權重就是tf.Variable。JAX 的數組默認“不可變”。每次運算返回新數組用jax.numpy替代 NumPy但 API 和 NumPy 高度一致遷移成本最低。張量形狀和類型轉換# PyTorch 查看和轉換 pt torch.randn(4, 8) print(pt.shape, pt.size()) pt_np pt.numpy() # 注意 requires_gradTrue 時不能直接轉換 pt2 torch.from_numpy(pt_np) # TensorFlow 查看和轉換 print(tf_t.shape) tf_np tf_t.numpy() # JAX 查看和轉換 print(jx.shape) jx_np jnp.asarray(jx)設備控制是三框架 API 差異最大的地方之一。PyTorch 使用.to(cuda)顯式搬運張量TensorFlow 在strategy.scope()或 Keras 里自動處理JAX 更徹底數據默認就在加速器上普通jnp運算會自動選擇可用設備。PyTorch 新手最常見的 Bug 就是把 CPU 張量直接傳給 CUDA 模型報 mismatch 錯誤。建議統一寫成model model.to(device)并把輸入輸出的設備邏輯封裝到訓練函數里。5. 自動微分 API 對比自動微分是深度學習框架的核心三者的實現思路完全不同。PyTorch 采用“動態計算圖 反向傳播”。張量開啟requires_gradTrue后前向執行時自動記錄梯度函數調用backward()后梯度回傳。代碼看起來和普通數值計算一致這是它容易上手的關鍵。import torch x torch.tensor([1.0, 2.0, 3.0], requires_gradTrue) y (x ** 2).sum() y.backward() print(PyTorch grad:, x.grad) # [2.0, 4.0, 6.0]TensorFlow 用tf.GradientTape的上下文管理器顯式記錄前向計算。import tensorflow as tf x tf.Variable([1.0, 2.0, 3.0]) with tf.GradientTape() as tape: y tf.reduce_sum(x ** 2) grad tape.gradient(y, x) print(TensorFlow grad:, grad.numpy())JAX 用純函數變換jax.grad沒有動態圖也沒有“反向傳播”這個動作而是直接對損失函數求梯度。要就求一階寫jax.grad要求“損失值和梯度一起拿”用jax.value_and_grad。import jax import jax.numpy as jnp def loss_func(x): return jnp.sum(x ** 2) x jnp.array([1.0, 2.0, 3.0]) grad jax.grad(loss_func)(x) print(JAX grad:, grad)JAX 的jax.jit編譯、jax.vmap向量化、jax.pmap多設備并行都是“函數變換”和 Python 原來的控制流不是一回事。寫 JAX 時要避免用純 Python 的if/for處理張量分支盡量用jnp.where、jax.lax.scan這類可變換結構否則編譯時機和性能都會有坑。6. 模型構建 API 對比PyTorch 用nn.Module。模型是類forward方法定義前向子模塊自動收集參數。import torch.nn as nn class MLP(nn.Module): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.fc1 nn.Linear(in_dim, hidden_dim) self.relu nn.ReLU() self.fc2 nn.Linear(hidden_dim, out_dim) def forward(self, x): return self.fc2(self.relu(self.fc1(x)))TensorFlow 高層接口是 Keras。Sequential適合順序結構Model適合多輸入多輸出和自定義結構。import tensorflow as tf model tf.keras.Sequential([ tf.keras.layers.Dense(hidden_dim, activationrelu), tf.keras.layers.Dense(out_dim) ])TensorFlow 還可以繼承tf.keras.Model寫自定義層和自定義前向風格上和 PyTorch 的nn.Module接近。團隊內部要想清楚“統一走 Keras 高層接口”還是“自定義代碼”兩種混用會讓維護成本上升。JAX 本身沒有內置模型類生態里最常用的是 Flax。Flax 用nn.Module但核心是“參數初始化函數 apply 方法”模型實例只是配置描述不保存權重。import flax.linen as nn import jax.numpy as jnp class MLP(nn.Module): hidden_dim: int out_dim: int nn.compact def __call__(self, x): x nn.Dense(self.hidden_dim)(x) x nn.relu(x) x nn.Dense(self.out_dim)(x) return x model MLP(hidden_dim64, out_dim1) params model.init(jax.random.PRNGKey(0), jnp.ones((1, 10))) pred model.apply(params, jnp.ones((1, 10)))Flax 的params是一個獨立字典訓練時通過參數傳遞更新。這種設計一開始會不習慣但配合optax做參數更新時非常清晰。JAX 生態里也有 Haiku、Equinox 等替代庫選型前先看團隊共識。7. 訓練循環 API 對比這是三個框架差異最明顯、也是遷移成本最高的部分。PyTorch 的訓練循環完全手寫邏輯全在你控制之下optimizer torch.optim.Adam(model.parameters(), lr1e-3) loss_fn torch.nn.MSELoss() for epoch in range(num_epochs): for x_batch, y_batch in dataloader: optimizer.zero_grad() pred model(x_batch) loss loss_fn(pred, y_batch) loss.backward() optimizer.step()TensorFlow 有兩種訓練方式。不想寫輪子就用model.fitmodel.compile(optimizeradam, lossmse) model.fit(x_train, y_train, epochs10, batch_size32, validation_split0.1)需要精細控制梯度時用GradientTape自定義訓練循環optimizer tf.keras.optimizers.Adam(learning_rate1e-3) loss_fn tf.keras.losses.MeanSquaredError() for epoch in range(num_epochs): for x_batch, y_batch in dataset: with tf.GradientTape() as tape: pred model(x_batch, trainingTrue) loss loss_fn(y_batch, pred) grads tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables))JAX 是“自己寫一切”但框架提供了可組合的變換。下面是一個最小訓練步配合optaximport optax import jax def loss_fn(params, x_batch, y_batch): pred model.apply(params, x_batch) pred pred.reshape(y_batch.shape) return jnp.mean((pred - y_batch) ** 2) optimizer optax.adam(learning_rate1e-3) opt_state optimizer.init(params) jax.jit def train_step(params, opt_state, x_batch, y_batch): loss, grads jax.value_and_grad(loss_fn)(params, x_batch, y_batch) updates, opt_state optimizer.update(grads, opt_state, params) params optax.apply_updates(params, updates) return params, opt_state, loss注意jax.jit裝飾后傳入的數據必須是數組而不是 Dataset 迭代器所以 JAX 的數據加載通常先取“一塊 numpy/tf.data 數據”再交給編譯后的train_step。這也是 JAX 和 PyTorch 訓練流程差異最大的地方。三者的選擇標準可以這樣記PyTorch 保留最大控制權且調試直接TensorFlow 的fit生產集成方便但自定義邏輯需要繞一下JAX 追求函數式純變換適合手寫科研算法但初始學習成本最高。8. 數據加載與預處理 API 對比數據管道在三大框架中各自獨立接口不通用。PyTorch 用DatasetDataLoader。自定義 Dataset 只需實現__len__和__getitem__DataLoader 自動負責 batch、打亂、多進程加載。from torch.utils.data import Dataset, DataLoader class MyDataset(Dataset): def __init__(self, x, y): self.x x self.y y def __len__(self): return len(self.x) def __getitem__(self, idx): return self.x[idx], self.y[idx] dataloader DataLoader(MyDataset(x_train, y_train), batch_size32, shuffleTrue, num_workers4)TensorFlow 用tf.data.Dataset。它的優勢是自帶管道優化prefetch、map、batch、cache都可以鏈式調用還能配合TFRecord做大規模數據流。import tensorflow as tf dataset tf.data.Dataset.from_tensor_slices((x_train, y_train)) dataset dataset.shuffle(buffer_size1000) dataset dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)JAX 沒有專用的 DataLoader社區常用tf.data或grainDeepMind 開源的數據加載庫。常見做法是先用tf.data.Dataset完成 map/batch/prefetch再用ds.as_numpy_iterator()喂給 JAX 訓練循環。注意 JAX 訓練循環通常用for batch in dataset:但batch是 numpy 數組直接傳給jitted函數即可。import tensorflow as tf dataset tf.data.Dataset.from_tensor_slices((x_train, y_train)) dataset dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE) for x_batch, y_batch in dataset.as_numpy_iterator(): params, opt_state, loss train_step(params, opt_state, x_batch, y_batch)從數據加載 API 來看PyTorch 靈活但多進程配置要調試TensorFlow 工程化強但 API 層級較多JAX 沒有標準答案靠組合。9. 部署與生態接口對比訓練結束后部署路徑決定了框架選型是否成功。PyTorch 常用導出方式是 TorchScript 和 ONNX。TorchScript 把模型編譯為可序列化圖適合 C 調用ONNX 是把模型遷移到其他運行時的重要通道很多加速卡廠商都支持 ONNX 導入。# PyTorch 導出 ONNX model.eval() dummy_input torch.randn(1, 10) torch.onnx.export(model, dummy_input, model.onnx, input_names[input], output_names[output])TensorFlow 部署是它的傳統強項。SavedModel 是標準格式配合 TF Serving 直接提供 gRPC/HTTP 推理服務TFLite 適合移動端、嵌入式TensorFlow.js 能跑在瀏覽器和 Node.js。模型轉換路徑清晰從訓練到生產不用切換體系。# TensorFlow 導出 SavedModel model.export(saved_model_dir)JAX 部署路徑相對“年輕”。常用方案是jax2tf把 JAX 函數轉成 TensorFlow 計算圖再做 SavedModel 導出也有團隊直接在生產環境用 XLA 編譯的 JAX 函數做推理服務。JAX 在分布式并行推理上能力很強但推理基礎設施需要自己搭沒有 TensorFlow Serving 那種開箱即用的組件。如果你打算把模型接到 API 平臺PyTorch 和 TensorFlow 都可以先用 ONNX/SavedModel 轉換再交給專用推理服務。JAX 則要提前驗證目標推理平臺是否支持 XLA 或jax2tf轉換否則部署環節會卡住。10. 資源占用與性能觀察框架本身不會直接告訴你“顯存夠不夠”要自己觀察和分析。通用觀察工具watch -n 1 nvidia-smiPyTorch 還可以在代碼里打印顯存分配print(fallocated: {torch.cuda.memory_allocated() / 1024**2:.1f} MB) print(freserved: {torch.cuda.memory_reserved() / 1024**2:.1f} MB)TensorFlow 默認會預占大量顯存調試時可以改成按需增長gpus tf.config.list_physical_devices(GPU) if gpus: tf.config.experimental.set_memory_growth(gpus[0], True)JAX 查看設備數量import jax print(jax.device_count())顯存占用主要受四個因素影響batch size、輸入尺寸分辨率/序列長度、模型參數量、優化器狀態。增大 batch 是訓練速度收益最明顯的手段但顯存壓力會同步上漲。如果顯存不足優先減 batch size而不是降分辨率。減 batch 還不行再用梯度累加模擬大 batch。混合精度是另一個常用手段PyTorch 用torch.cuda.ampTensorFlow 用mixed_float16JAX 配合jax.disable_float32()或顯式使用bfloat16。CPU 和 GPU 的性能差距很難給統一數字因為運算類型、數據量、環境都不同。要觀察就固定數據規模分別跑 20 個 batch記錄耗時和顯存變化。對比時用同一套超參不要一邊帶編譯優化一邊不帶否則對比結果沒有參考價值。11. 常見問題與排查方法問題現象可能原因排查方式解決方案PyTorch 裝完torch.cuda.is_available()為 FalseCUDA 版本與 PyTorch wheel 不匹配nvidia-smi看驅動檢查安裝命令 index-url去 pytorch.org 重新生成匹配 CUDA 的安裝命令TensorFlow 2.18 裝完檢測不到 GPU缺少 CUDA/cuDNN 運行庫檢查是否安裝tensorflow[and-cuda]Linux 安裝tensorflow[and-cuda]Windows 對照官方文檔配 CUDA DLLJAXjax.devices()只顯示 CPUJAX CUDA 版未安裝或庫沒配對打印jax.__version__確認安裝的是jax[cuda12]等 GPU 包用對應 CUDA 版本的jax[cuda12]/jax[cuda11]重新安裝訓練時顯存不足 OOMbatch size 過大、輸入尺寸過大、優化器狀態過多nvidia-smi看占用峰值減小 batch size使用梯度累加、混合精度、gradient checkpointingDataLoader 多進程卡死PyTorch Windows 下num_workers配置不當把num_workers調為 0 測試將啟動代碼放入if __name__ __main__:按系統調整num_workersPyTorch 加載老模型報錯PyTorch 2.6 起torch.load默認weights_onlyTrue檢查加載代碼手動指定weights_onlyTrue或對可信權重使用完整加載并明確處理反序列化風險JAX 訓練在 CPU/GPU 之間跳數據不是數組進入了不支持變換的 Python 控制流打印訓練輸入類型統一用jnp數組避免在jit裝飾函數里用 Python 原生if/for判斷張量model.fit效果正常自定義 GradientTape 報錯變量沒有用tf.Variable包裝檢查模型參數是否在trainable_variables自定義層和模型都繼承tf.keras.Model讓框架托管參數排查原則是“先環境后代碼”。報錯先看驅動、CUDA、Python 版本是否匹配再看數據形狀和設備是否一致。三套框架的報錯信息里都會給出設備、張量形狀和具體操作位置不要只看第一行。12. 最佳實踐與使用建議工程上不管是 PyTorch、TensorFlow 還是 JAX以下做法都適用。第一第一次跑新環境先小規模驗證。不要直接上完整模型和大 batch先跑 1 個 batch、10 步訓練確認前向、反向、優化器、保存全部能通再擴規模。這樣能快速區分“環境問題”和“算法問題”。第二環境隔離和版本固定。用 conda 或 venv 為每個項目建立獨立環境要求項目里記錄 Python、框架、CUDA、關鍵依賴的精確版本。框架升級帶來的兼容性問題比多數模型本身的問題更難排查。第三目錄規范。建議按data/、models/、outputs/、src/分層管理原始數據、模型權重、日志、訓練腳本分開。JAX 和 PyTorch 的模型權重格式不通用直接拆目錄存params.pt、saved_model、params.pkl避免一個目錄堆滿二進制文件。第四批量訓練任務要加日志和斷點。PyTorch 訓練循環里加torch.save斷點非常自然TensorFlow 用ModelCheckpoint回調JAX 需要自己把params和opt_state序列化。沒有斷點機制就大規模訓練任何一個節點中斷都會浪費大量算力。第五模型保存和加載要適配版本。PyTorch 2.6 起torch.load默認weights_onlyTrue加載舊權重時先確認反序列化安全性。TensorFlow 的 SavedModel 和 Keras.h5格式不要混用JAX 生態的權重一般配合 Flax/optax 結構保存。最后合規提醒要前置。訓練數據、人臉數據、語音數據、版權素材都要確認授權模型部署為 API 服務時要加訪問控制對外不能無鑒權裸奔使用開源模型權重先檢查許可證。做技術驗證沒問題公開上線或商用前必須走法務和合規檢查。13. 總結與下一步這篇文章的核心結論是三條路對應三種思維方式PyTorch 讓你像寫普通 Python 一樣寫模型和訓練邏輯調試成本最低TensorFlow 給你從訓練到部署的最完整工程鏈路適合團隊標準化交付JAX 讓你用函數變換組合出高效計算流程適合科研算法和需要極致并行控制的項目。三者的核心 API 差異集中在張量可變性、自動微分方式、模型組織形式、訓練循環控制權和數據加載方式五個維度。建議你拿到任何新框架都先跑一遍同樣的最小測試安裝驗證、張量創建、求梯度、搭一個兩層的 MLP、寫一個 10 步訓練循環、導出模型。誰都能用這套測試在半小時內跑通整個鏈路。最容易踩的坑是“用 PyTorch 的思維寫 JAX”或者“用 TensorFlow 的高層接口寫完卻想在自定義訓練循環里接管所有參數”。框架之間的遷移不是改 API 名而是改代碼的組織方式。下一步可以往三個方向擴展一是對比三者的分布式訓練接口DistributedDataParallel、tf.distribute.Strategy、jax.pmap完全是三種抽象二是研究 ONNX 作為跨框架交換格式把 PyTorch 和 TensorFlow 模型統一部署三是深入 JAX 的vmap/jit組合它的性能上限和調試復雜度都值得單獨寫一篇。建議先把最小測試在三個框架上都跑通后續再按任務類型選型別急著在大項目里做一次性遷移。