从单卡到集群(多台连成一群一起干活的机器)再到生产:模型怎么走出笔记本
这一章讲两件事: 训练怎么摊到多张卡、多台机器上(原书第 16 章);训完的模型怎么变成别人能用的服务(原书第 17 章)。 位置:全书收尾。前面所有章都假设「一张卡、一个进程、自己跑」;这一章把两个假设都撤掉。
1. 为什么一张卡不够
第 13 章用 DataParallel 蹭过多卡,但那只在单机里打转。真正的动力来自规模:书里给了个坐标——Meta 的 LLaMA 3.1 最大变体(405B 参数)光是推理就要约 800GB 显存(显存 = GPU 上的内存,模型权重和中间结果都得住进去),而一张顶级 GPU 的显存只有它的十分之一上下1。训练要的更多:除了权重,还要放梯度和优化器状态。
并行(并行计算,把一份活拆给多个工人同时干)因此不是优化,是入场券。这章前半就是讲:活怎么拆,拆了怎么不出错。
2. 分布式的词汇表
分布式训练(distributed training)是让多个进程一起完成一次训练。为什么用进程不用线程:Python 的线程共享内存、互相踩脚;进程各自隔离,各管一张 GPU,靠消息通信——隔离换来了干净2。
五个词先认下,后面全靠它们说话2:
- world size:进程总数。一个进程管一张卡,所以通常等于 GPU 数。
- rank:每个进程的编号,0 到 world size−1,像工号。
- process group:能互相通信的一组进程,初始化时划定。
- node:一台物理机器;一台机器里可以跑多个进程。
- SPMD(单程序多数据):所有进程跑同一份代码,各自用 rank 判断「我该干哪份」——分布式程序的默认写法。
起跑要三步:拉起进程 → 让进程互相找到对方 → 改训练逻辑。第二步靠一个「会合点」:所有进程连到同一个 TCPStore(一个键值存储,地址由 MASTER_ADDR/MASTER_PORT 这两个环境变量(操作系统层面传给程序的键值对)指定),报到、领号。手动做很啰嗦,所以 PyTorch 给了 torchrun:它读环境变量替你拉起全部进程,你的代码只管后两步3。
3. 集合通信:broadcast 与 allreduce
进程间的标准动作叫集合通信(collective communication),本章只要两个4:
- broadcast:一个进程发、所有人收。书里比作电台:rank 0 是电台,其余是听众。用途:把同一份初始权重发给所有进程,保证起跑线一致。
- allreduce:所有人各出一份数据,按某种运算(求和、取大…)合并,结果再发回所有人。书里比作凑菜:每人带一样食材,最后每人都分到一盘完整的菜。
跑一次(书里 4 进程的例子):rank 0–3 各持一个张量 t0–t3,allreduce(求和)之后,四个进程手里都是 t0+t1+t2+t3。这一步是全章的机械核心,下一节马上用到4。
4. 主走查:数据并行的一次迭代
最常用的拆法是数据并行(data parallelism):模型每个进程一份完整拷贝,数据切开,各算各的。走查 3 张卡的一次迭代,每步旁边是状态5:
起跑:broadcast 后,三张卡权重完全相同(工号 0/1/2,world size=3)
① 取数据:DistributedSampler 给每张卡分一批不重复的数据
rank0 拿到样本 0–7,rank1 拿 8–15,rank2 拿 16–23
② 前向+反向:各算各的损失与梯度
rank0 的梯度 g0,rank1 的 g1,rank2 的 g2 —— 三个数不一样!
③ 同步(关键):allreduce(求和)→ 每张卡得到 g0+g1+g2,再除以 3
三张卡的梯度重新一致
④ 更新:各自 optimizer.step() —— 输入相同,结果也相同,权重保持同步
图说:前向天然能拆,麻烦全在第 ③ 步:不齐一次梯度,模型就分叉。
第 ② 步的天然并行有个名字:前向计算是「embarrassingly parallel」(易并行,各批数据互不依赖,拆就完了)5。真正的坑在第 ③ 步:三张卡见的数据不同,梯度自然不同——如果跳过同步直接 step,三份权重从此走上三条路,等于白训。allreduce 把梯度拉回一致,是数据并行里唯一不可省的开销5。
手写这几步能跑,但生产里直接用 DDP(DistributedDataParallel,第 13 章预告过):包住模型,它在反向传播时自动在后台做梯度 allreduce,还懂得和计算重叠省时间。配套的就是第 ① 步的 DistributedSampler,保证每卡每轮拿到的数据片不重复5。
5. 模型装不下时:另外三种并行与 FSDP
数据并行有个硬前提:模型得塞得进单卡。塞不下时,按模型的「身材」选拆法6:
- 模型并行:把模型的不同层放到不同卡上,中 间结果用点对点 send/recv 传递。缺点赤裸:同一时刻只有一张卡在算,其余干等。
- 流水线并行:给模型并行配上微批次(microbatch,把一批数据再切小):卡 A 算微批次 1 时,卡 B 在算微批次 0——像流水线,各站同时忙。书里用的 GPipe 排程先把所有微批次的前向跑完,再统一跑反向。适合「又深又瘦」的模型。
- 张量并行:在一层的内部切参数矩阵(按行或按列),一层都算不完时用。通信量大,通常只在一台机器内部搞。适合「又宽又浅」的模型。
真实大模型是组合拳:书里提到 LLaMA 3 用了 4 维并行(数据+张量+流水线+上下文)。管这套组合的工具是 DeviceMesh:把设备声明成网格(比如 2×4),按维度名自动建好各组通信通道6。
还有一条更激进的路:FSDP(全分片数据并行)。数据并行每卡存一份完整参数,800GB 的模型免谈;FSDP 干脆把参数本身也切片,每卡只存 1/N,算到某层时临时 all_gather(全员凑齐该层),用完即弃,反向用 reduce_scatter 分摊梯度。显存换通信,512 卡规模验证过;模型大到 DDP 不行时,书里推荐它7。
大语言模型(LLM,用海量文本训出来的语言模型)还催生两种专用拆法,一句话各归其位:上下文并行把超长序列沿长度切给多卡(注意力机制的内存随长度平方涨);专家并行服务专家混合(MoE——Mixture of Experts,把前馈层换成一组「专家」小网络,每个 token 只送路由选中的几个)——专家散在多卡,靠全对全通信送数据8。想动手,官方参考实现 TorchTitan 收好了 LLaMA 等模型的现成配置9。
6. 部署:从能跑到能用
训练完,模型还是你笔记本里的一个对象。后半章走五条把它交出去的路,按离生产的距离排10:
Gradio 出原型。 三行代码(gr.Interface(函数, 输入控件, 输出控件))自动生成网页界面,share=True 还能给个公网链接给同事试。定位明确:原型工具,不是生产11。
FastAPI 出服务。 把模型装进一个 HTTP 服务:用 Pydantic 模型声明输入格式(自动校验、自动生成 /docs 文档页),端点函数里跑模型。书里的例子是 SmolLM2-360M 这个小语言模型12。
接下来是两个真正「生产级」的动作,先各讲清楚,再合起来走一遍13。
动作一:别让端点干等。 跑这个服务的那台机器叫服务器——常年开着,随时等着接请求。最朴素的写法是接一个、算完、回一个,串着来;这样 GPU 大部分时间都在干等 Python 处理杂活。
在端点函数前面加一个 async,这个端点就变成异步——不等手上这件事做完,就先去做别的。等 GPU 出结果的那段时间里,这条线程可以掉头去接下一个请求。
手上同时压着好几件事、谁先好了就先处理谁,这个能力叫并发。它的瓶颈是 Python 的 GIL(全局解释器锁:同一时刻只有一个线程在执行 Python 字节码),所以真正的重活必须交给后台线程去干,不能压在端点函数里。
动作二,请求批处理。 用户的请求一个一个来,GPU 却喜欢成批干活。解法是解耦:端点只把请求丢进队列,后台 worker 凑满一批(或者等到超时)再一次跑完、逐一分发结果。书里比作餐厅:点菜单攒成一叠,厨师一次开炒。流式输出是同一个结构的副产品:每生成一个词就立刻推回给用户,不用等整段生成完。
部署这一半的走查:一个请求从进门到吐出第一个字
上面这些零件各是什么,已经点清楚了。但真正要答的那个问题还没答:一个请求 打进来之后,到底依次发生了什么。 所以拿书里那个 SmolLM2-360M(3.6 亿参数的小语言模型)走一遍,从用户敲下回车开始1213:
① 进门与校验
POST /generate body: {"text": "从前有座山"}
Pydantic 声明的 TextInput 只有一个字段 text: str。
合格 → 拿到那个字符串;
不合格(传成 {"txt": …} 或 {"text": 123})→ 框架当场回 422,
模型一次都没被碰到。这就是「自动校验」买到的东西。
② 发号、入队,函数立刻返回
生成 request_id = "a3f1…"(一个 uuid),
把 (request_id, "从前有座山") 丢进 inference_queue,然后 return。
端点函数**没有等模型**。正因为它是 async 的,这条线程马上能去接下一个请求。
③ 后台 worker 攒批
一个常驻线程死循环地从队列里捞:凑够 8 个就走,凑不够最多等 50 毫秒。
如果第 3 个请求进来时离第 1 个已经过了 50 毫秒,就以 3 个成批,不再干等。
—— 这一步是延迟和吞吐在打架:等得久,批更大、GPU 更划算,
但先到的那个请求白白多等了 50 毫秒。
④ 一次前向 = 这一批里每人各出一个字
model.generate(…, max_new_tokens=1) ← 注意是 1,不是 256
8 个提示一起送进 GPU,一次前向出 8 个 token,一人一个。
设这次前向花 40 毫秒:8 个请求分摊,人均 5 毫秒。
同样这 8 个请求要是一个一个跑,是 8 × 40 = 320 毫秒。
⑤ 逐字推回去
每个新 token 立刻塞进它自己 request_id 名下的结果队列;
StreamingResponse 从那个队列里一取到就往客户端推。
所以用户在第 ④ 步结束时(约 90 毫秒)就看见了第一个字,
而不是等 256 个字全生成完(256 × 40 毫秒 ≈ 10 秒)才看到东西。
⑥ 回到第 ④ 步,再来一轮
每轮只出一个字,一圈一圈转,直到这一批里所有请求都吐完或者撞上字数上限。
——你在聊天界面里看到字一个个往外蹦,就是这个循环。
图说:批大小 8、超时 50 毫秒、单次前向 40 毫秒这三个数是为演示编的
(书里的代码把它们收在常数里,正文没印具体值);
「max_new_tokens=1」「uuid 发号」「丢队列后立刻返回」「StreamingResponse 逐词推」
都是书里代码的原样做法。
这条走查上没有 ONNX、torch.export、LibTorch。 它们不是这条路上的某一步,而是一个分岔口,岔在第 ④ 步之前:如果你压根不打算用 Python 跑这个模型,就在上线之前把它导成别的格式,然后整条 ①–⑥ 换一门语言重写。下一节的三件套改的是第 ④ 步本身;再下一节讲这三条岔路各通向哪儿。
7. 让模型变小变快
服务上线后,下一仗是成本。走查里唯一真正花时间的是第 ④ 步那次前向(40 毫秒),所以这一节的三件套全都只动那一步——它们不改流程,只让第 ④ 步更便宜14:
- torch.compile:PyTorch 2.0 的即时编译开关,包一行
torch.compile(model),把一串小操作融合成大内核、重排计算,推理提速 20%–200%(偶尔也更慢,要实测)。 - 蒸馏(distillation):训一个小模型去模仿大模型的输出,拿体积换精度——DistilBERT 就是这么来的。
- 量化(quantization):把权重从 32 位浮点改成 8 位整数存储和计算。书里走查:SmolLM 从 6528MB 压到 1825MB,省 70%。不能整模型一刀切(归一化层、激活还要浮点),所以用
quantize_dynamic只量化 Linear 层。为什么信息砍半还能用:舍入误差近似随机噪声,而卷积和线性层做的是加权平均——误差在求和里互相抵消14。
8. 导出与出 Python
最后三条路是走查里说的那个分岔口:都是「离开 PyTorch 的 Python 解释器」,一旦走上去,前面那条 ①–⑥ 的服务链就要在别的运行时里重搭一遍15。
- ONNX:开放的模型交换格式。
torch.onnx.export(模型, 示例输入, 路径)跟踪一次前向、把计算图序列化成文件(序列化 = 把内存对象落成字节流),任何兼容运行时都能跑——树莓派上比直接跑 PyTorch 还快。注意是单行道:PyTorch 不能读回 ONNX。 - torch.export + AOTInductor:PyTorch 自家的提前编译链,产出
.pt2包,C++ 环境直接加载。和 compile 的分工:export 要求整图捕获,遇到「运行时才定走向」的动态控制流(比如if x.sum()>0这种依赖数据的分支) 直接报错;torch.compile 则宽容,捕获不了的地方「断图」回退到普通 Python 接着跑。 - LibTorch / ExecuTorch:前者是 C++ 版 PyTorch(API 几乎是 Python 的镜像),写给不能带 Python 的服务端;后者面向手机和边缘设备(手机、嵌入式板子这类远离机房的设备),
torch.export → to_edge → to_executorch一条链,绕开 JNI 那层麻烦。
优化有没有用,别凭感觉:torch.profiler 包住一次前向,按操作名列耗时——书里的例子中矩阵乘 aten::addmm 占了 65% 的 CPU 时间,优化矛头该指哪一目了然;导出的 trace 还能对比编译前后的执行面貌16。
9. 作者的判断与证据
| 说法 | 性质 |
|---|---|
| 405B 推理约 800GB 显存 | 有据,粗算可复(参数量×每参数字节)1 |
| 梯度不同步模型会分叉 | 机制必然,数学上可证5 |
| 流水线/张量并行按模型身材选 | 作者经验法则,方向可靠6 |
| 请求批处理提吞吐(单位时间能处理多少请求) | 有据,GPU 批处理原理;排队会加延迟(一个请求要等多久),书里也给了流式补偿13 |
| compile 提速 20%–200% | 作者实测口径,明说「偶有更慢」14 |
| 量化 6528→1825MB 且精度可用 | 有据(单个模型实测);「误差相互抵消」是作者的机制解释,方向对14 |
10. 边界与局限
- 这章的并行代码都是教学版(gloo/CPU 后端);真上 GPU 集群用 nccl,细节(网络拓扑、故障恢复)远超本章4。
- FSDP 省显存但通信量上升,慢网络下可能得不偿失——书里给了选型方向,没给量化判据7。
- 请求批处理用延迟换吞吐:实时性要求 高的场景(对话)要靠流式找回体验13。
- ONNX 是单行道且算子覆盖有限;torch.export 要求输入形状与导出时一致——书里建议永远保留能从源码重新导出的工作流(一套固定走下来的流程)15。
- 对抗样本(故意构造的恶意输入)在部署章只提了一句没展开——上线系统要另查。
11. 可带走的
- 并行先问「瓶颈在哪」:数据多 → 数据并行;模型大 → 按身材选流水线(深瘦)或张量(宽浅);参数都放不下 → FSDP。
- 数据并行的命门是梯度同步:allreduce 每步一次,省不得;手写一遍再交给 DDP。
- 分布式词汇就五个:world size、rank、process group、node、SPMD。
- 部署按距离选工具:Gradio 原型、FastAPI 服务、请求批处理喂饱 GPU、流式补体验。
- 一个请求的完整旅程记得住:校验 → 发号入队、端点立刻返回 → 后台凑够一批或等到超时 → 一次前向、批里每人出一个字 → 逐字流回 → 再转一圈。攒批换的是吞吐,赔的是先到者的延迟;流式把这笔账在体感上补回来。
- 成本三件套只动「一次前向」那一步:compile(编译)、蒸馏(换人)、量化(降精度)— —量化优先,一行 API 省 70%。
- 导出是单行道,而且是岔路不是台阶:走 ONNX / export / LibTorch 就要在别的运行时重搭整条服务链;保留从源码重新导出的能力;动态控制流用 compile,不用 export。
- 优化前后都跑 profiler,让耗时表代替直觉。
12. 原文地图
| 主题 | 原书章 | 原文位置 |
|---|---|---|
| 800GB 显存 | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:14(搜「LLaMA 3.1」) |
| 进程 vs 线程 | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:98(搜「world size」) |
| torchrun/TCPStore | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:218(搜「TCPStore」) · :238(搜「torchrun」) |
| broadcast/allreduce | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:291(搜「radio station」) · :296(搜「contributing ingredients」) |
| 数据并行主走查 | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:376(搜「embarrassingly parallel」) · :492(搜「wrapping our model with DDP」) · :506(搜「DistributedSampler」) |
| 三种并行选型 | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:533(搜「model parallelism」) · :613(搜「microbatches」) · :732(搜「skinnier and deeper」) |
| DeviceMesh/4D | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:764(搜「4D parallelism」) · :802(搜「init_device_mesh」) |
| FSDP | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:841(搜「FSDP」) · :853(搜「all_gather」) |
| MoE 与上下文并行 | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:931(搜「Mixture of Experts」) |
| TorchTitan | ch16 | text/24-ch16-16-training-models-on-multiple-gpus.txt:953(搜「TorchTitan」) |
| Gradio | ch17 | text/25-ch17-17-deploying-to-production.txt:61(搜「Gradio is a Python library」) · :137(搜「share=True」) |
| FastAPI/Pydantic | ch17 | text/25-ch17-17-deploying-to-production.txt:147(搜「FastAPI is a popular」) · :184(搜「Pydantic model」) |
| 异步/批处理/流式 | ch17 | text/25-ch17-17-deploying-to-production.txt:277(搜「asynchronous」) · :319(搜「Request batching」) · :357(搜「busy restaurant」) |
| GIL | ch17 | text/25-ch17-17-deploying-to-production.txt:701(搜「global lock」) |
| compile/蒸馏/量化 | ch17 | text/25-ch17-17-deploying-to-production.txt:515(搜「PyTorch 2.0」) · :559(搜「distillation」) · :598(搜「6528.52 MB」) · :658(搜「1825.69 MB」) |
| ONNX | ch17 | text/25-ch17-17-deploying-to-production.txt:721(搜「ONNX」) · :751(搜「torch.onnx.export」) |
| export/断图/Dynamo | ch17 | text/25-ch17-17-deploying-to-production.txt:782(搜「traced graph」) · :989(搜「Dynamic control flow is not supported」) · :1018(搜 「torch.jit.trace」) |
| profiler | ch17 | text/25-ch17-17-deploying-to-production.txt:1088(搜「torch.profiler」) · :1130(搜「aten::addmm」) |
| LibTorch/ExecuTorch | ch17 | text/25-ch17-17-deploying-to-production.txt:1179(搜「LibTorch」) · :1320(搜「to_edge」) |