外观
框架与工具怎么选
一句话定义:深度学习框架是"研究速度 × 工程效率 × 生态兼容"的三维权衡,没有绝对最优,只有与你的任务、团队、基础设施最匹配的选项。
本文先横向对比三大主流框架,再纵向梳理训练基础设施、推理框架与生态工具,最后给出决策树。上手路径可参考学习路径:三条路线;本文默认你已理解训练的基本流程(见从零构建一个深度学习项目)。
一、三大框架横向对比
| 维度 | PyTorch | JAX | TensorFlow / Keras |
|---|---|---|---|
| 图范式 | 动态图为主(eager),支持 torch.compile 编译 | 函数式纯变换(jit 即时编译到 XLA) | 默认静态图,Keras 提供"类 eager"接口 |
| 易用性 | ★★★★★ 直觉、易调试 | ★★★ 函数式 + 纯函数,学习曲线陡 | ★★★ 早期混乱,2.x 后改善 |
| 调试 | print/断点/pdb 直接可用 | 不易打印中间量(需 jax.debug 钩子) | tf.debugging 工具,Keras 下较友好 |
| 性能 | 好;torch.compile/inductor 进一步优化 | 强(XLA 编译、jit、pmap/shard_map 大规模并行) | 强;部署链路成熟 |
| 分布式 | DDP/FSDP 成熟,PyTorch 2 原生 | pmap/pjit/shard_map 抽象更底层更灵活 | tf.distribute 生态完整 |
| 生态 | 研究界事实标准,HuggingFace 全系支持 | 谷歌系 + DeepMind 研究、TPU、大模型缩放研究 | 工业界存量、Keras 快速原型、移动端 |
| 社区/招聘 | 当前最主流 | 上升期 | 存量最大但增速放缓 |
一句话总结:做研究、读论文复现、跟 HuggingFace 生态 → PyTorch;做大规模并行/TPU、函数式科学计算、研究缩放定律 → JAX;已有 TF 存量系统、快速原型、跨平台部署(移动端)→ TensorFlow/Keras。
一个务实提醒
现在新项目用 TensorFlow 的越来越少,但公司存量代码里 TF 仍然很多。面试和工作中"能看懂 TF 代码"仍是加分项(见JD 清单),但新项目优先考虑 PyTorch 或 JAX。
二、训练基础设施:从单卡到多卡
1. 单卡
一切起点。做好三件事:显存管理(混合精度、梯度累积,见训练配方与调参)、设备无关的代码(device = "cuda" if torch.cuda.is_available() else "cpu")、监控(nvidia-smi、W&B/TensorBoard)。
2. 多卡:DDP 与 FSDP
| 方案 | 用途 | 机制 | 适用规模 |
|---|---|---|---|
| DataParallel(已不推荐) | 单机多卡入门 | 每 batch 在每卡复制一份模型、梯度回传主卡 | <4 卡 |
| DistributedDataParallel(DDP) | 单机/多机多卡标配 | 每进程一份模型,梯度 all-reduce 同步 | 8–64 卡 |
| FullyShardedDataParallel(FSDP) | 大模型(微调/预训练) | 参数、梯度、优化器状态分片到各卡 | 单卡装不下的模型 |
| 张量并行/流水并行 | 超大规模 | 单层内切分/按层切分流水 | >64 卡、百亿级 |
PyTorch DDP 最小示例:
python
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def init(rank, world_size):
dist.init_process_group("nccl", rank=rank, world_size=world_size)
torch.cuda.set_device(rank)
model = DDP(YourModel().to(rank))
return model
if __name__ == "__main__":
import torch.multiprocessing as mp
mp.spawn(init, args=(torch.cuda.device_count(),), nprocs=torch.cuda.device_count())启动命令:torchrun --nproc_per_node=4 train.py。
FSDP 的关键经验:sharding_strategy 的选择、CPU offload 的开销、与混合精度/梯度检查点叠加时的显存收益。FSDP 是 HuggingFace 大模型微调(LoRA 等)的默认搭配之一。
3. 云平台与实验管理
- GPU 云:按小时租卡做实验(各云厂商 GPU 实例);训练任务管理用
accelerate/torchrun。 - 实验追踪:W&B、Neptune、MLflow 记录指标与配置;本地先用 TensorBoard +
config.yaml(见DL 设计原则)。 - 规模化训练:Kubernetes + 任务编排是 MLOps 范畴,见MLOps 与模型部署。
三、推理框架:把模型变成产品
训练完的模型要上线,推理链路与训练链路是两套工具:
| 框架 | 定位 | 特点 | 适用 |
|---|---|---|---|
| ONNX Runtime | 跨框架互操作标准 | 从 PyTorch/TF 导出 ONNX,CPU/GPU 都行 | 中规中矩、兼容性最强 |
| TensorRT | NVIDIA GPU 专属 | 层融合、量化(FP16/INT8),延迟最低 | 高吞吐 GPU 推理 |
| OpenVINO | Intel 硬件 | CPU/GPU/NPU 优化 | 边缘端 Intel 设备 |
| TVM / Apache TVM | 编译式推理栈 | 自动调优(AutoTVM)、多硬件后端 | 异质硬件、自定义算子 |
| vLLM | LLM 专用推理引擎 | PagedAttention、连续批处理、高吞吐 | 自托管大模型推理 |
选型逻辑:
- 通用模型:PyTorch → 导出 ONNX → ONNX Runtime,成本最低。
- 极致性能(单卡 GPU):转 TensorRT,典型获得 2–5 倍提速与更低延迟。
- 大语言模型服务:直接用 vLLM 或 TGI,别自己写推理脚本(KV cache、批处理策略见大语言模型(LLM))。
- 边缘设备:考虑 TensorRT/OpenVINO 的量化模型(INT8),或直接用 TFLite(如果走 TF 生态)。
bash
# 最小 ONNX 导出示例
torch.onnx.export(model, dummy_input, "model.onnx",
input_names=["input"], output_names=["output"],
dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}})四、HuggingFace 生态:现代深度学习的"标准库"
HuggingFace(HF)已经事实上成为模型、数据集与工具的标准分发渠道。四个核心库:
| 库 | 作用 | 典型用法 |
|---|---|---|
transformers | 预训练模型统一接口 | 加载/微调 BERT、GPT、ViT、Whisper…… |
datasets | 数据集统一加载与处理 | 流式加载大数据集、map 预处理、多进程缓存 |
diffusers | 扩散模型工具链 | 文生图、LoRA 微调 Stable Diffusion(见扩散模型与生成式 AI) |
peft | 参数高效微调 | LoRA/QLoRA 微调,一张消费卡微调 7B 模型 |
python
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.1-8B")
lora = LoraConfig(r=8, lora_alpha=16, task_type="CAUSAL_LM")
model = get_peft_model(model, lora) # 只训练 LoRA 参数,可训练参数量骤降
# datasets 示例:流式读取并批量预处理
from datasets import load_dataset
ds = load_dataset("imdb", split="train").select(range(1000))为什么值得学 HF 生态:它把"下载模型/数据、微调、评估、发布"压缩成了统一 API,是研究复现、比赛、作品集项目的加速器;peft 让大模型微调在消费级 GPU 上成为可能(见作品集项目)。更多资源见精选资源清单。
五、决策树:我怎么选?
按你的场景走一条路径:
你的首要目标是什么?
├─ 学深度学习/复现论文/做研究
│ └─ PyTorch(HF 生态加持)──────────→ PyTorch + transformers
├─ 大规模模型预训练/科学计算/TPU
│ └─ JAX(函数式、XLA、pmap)────────→ JAX
├─ 公司存量系统/移动端部署/Keras 快速原型
│ └─ TensorFlow/Keras ───────────────→ TF2 + TFLite
└─ 把已有模型上线
├─ 通用服务 → ONNX Runtime
├─ GPU 极致性能 → TensorRT
├─ LLM 服务 → vLLM
└─ 边缘端 → OpenVINO / TFLite
训练规模?
├─ 单卡装得下 → 混合精度 + 梯度累积(不用分布式)
├─ 单机多卡 → DDP
└─ 单卡装不下 → FSDP / LoRA(HF peft)→ 必要时再上分布式组合示例:一个典型的 2026 年个人项目栈 = PyTorch + HuggingFace(transformers/datasets/peft)+ ONNX Runtime(部署)+ W&B 或 TensorBoard(追踪),多卡时加 DDP/FSDP。
六、权衡与边界
框架不是永恒的。三个提醒:
- 别把框架当信仰:PyTorch 曾是"动态图"的代表,如今
torch.compile把静态优化能力也收了进来;JAX 也提供了更多 eager 调试手段——边界在持续移动。 - 换框架成本很高:项目与团队积累的生态绑定(模型库、算子、工具链)远超"重写一遍训练循环"。除非有硬性收益(性能、部署、团队技能),否则不换。
- 框架之争不如工具链之争:实际影响你产出的往往是生态工具——HF、W&B、vLLM——是否在你的框架上支持得最好。选框架前,先确认你需要的生态工具链是否齐全。
延伸阅读
- 学习路径:三条路线——按职业方向选技术栈
- MLOps 与模型部署——从模型到产品的完整链路
- 从零构建一个深度学习项目——PyTorch 最小项目的载体
- 大语言模型(LLM)——vLLM/peft 背后的模型知识
- 精选资源清单——各框架官方教程索引
- JD 清单——看看目标岗位实际要求哪个框架
参考资料
- PyTorch. Distributed data parallel——DDP 官方文档
- JAX. JAX: Autograd and XLA——JAX 官方文档
- Hugging Face. Transformers Documentation——HF 生态官方文档
- ONNX Runtime. Official Documentation——跨框架推理标准