如何扩展你的模型(8):在 TPU 上部署 LLaMA 3
中英对照第 8 篇:在 TPU v5e 上服务 LLaMA 3 要花多少钱?吞吐和延迟如何权衡?左栏原文,右栏译文。
本篇属于系列 如何扩展你的模型(How To Scale Your Model) · 第 8 篇
原文:How To Scale Your Model(Google DeepMind,作者 Jacob Austin、Sholto Douglas、Roy Frostig、Anselm Levskaya、Charlie Chen、Sharad Vikram、Federico Lebron、Peter Choy、Vinay Ramasesh、Albert Webson、Reiner Pope)。
本站是中英对照排版:左栏是英文原文,右栏是对应的中文译文。不好直译的术语(roofline、strong scaling、ICI、MXU、KV cache、FSDP、TP 等)保留英文写法;原文配图全部保留。
本文是系列《如何扩展你的模型(How To Scale Your Model)》的第 8 篇,共 13 篇。系列目录。
原文以 MIT 许可证发布,版权归 Google LLC;本译文仅作学习交流之用,如有错漏以原文为准。
This section will look at what it takes to serve LLaMA-3 and how efficiently it can be done. As in the previous “applied” section, try to work out the answers on your own with a pen and paper before looking them up!
本节会看服务 LLaMA-3 需要什么、以及能做到多高效。和上一个「实战」小节一样,先拿笔和纸自己试着算答案,再去看参考!
What’s the LLaMA Serving Story?
LLaMA 服务的故事是怎样的?
Let’s remind ourselves what LLaMA 3-70B looks like (see Section 6 for reference):
先回顾一下 LLaMA 3-70B 长什么样(可参考第 6 节):
| hyperparam | value |
|---|---|
| $$n_\text{layers}$$ (L) | 80 |
| $$d_\text{model}$$ (D) | 8,192 |
| $$d_{ff}$$ (F) | 28,672 |
| $$n_\text{heads}$$ (N) | 64 |
| $$n_\text{kv heads}$$ (K) | 8 |
| $$d_\text{qkv}$$ (H) | 128 |
| $$n_\text{embeddings}$$ (V) | 128,256 |
| 超参数 | 取值 |
|---|---|
| $$n_\text{layers}$$ (L) | 80 |
| $$d_\text{model}$$ (D) | 8,192 |
| $$d_{ff}$$ (F) | 28,672 |
| $$n_\text{heads}$$ (N) | 64 |
| $$n_\text{kv heads}$$ (K) | 8 |
| $$d_\text{qkv}$$ (H) | 128 |
| $$n_\text{embeddings}$$ (V) | 128,256 |
Let’s start with a simple question: what hardware should we serve on? The answer is basically, whichever is cheapest in FLOPs / dollar. (This isn’t always true, sometimes more HBM or ICI bandwidth is critical rather than FLOPs, but this is a good heuristic.) For this reason, we typically want to serve on TPU v5e, our current dedicated inference chip (cost comes from Google Cloud pricing as of February 2025):
先从一个简单问题开始:我们应该在什么硬件上服务? 答案基本上是:FLOPs / 美元最便宜的那个。(这并不总是成立,有时更关键的是 HBM 或 ICI 带宽而不是 FLOPs,但这是个很好的启发式。)因此,我们通常想用当前专用的推理芯片 TPU v5e(价格来自 Google Cloud 定价,截至 2025 年 2 月):
| TPU type | bfloat16 FLOPs/s | Google Cloud USD / hour | FLOPs / $ |
|---|---|---|---|
| H100 | 9.9e14 | $10.8 | 3.3e17 |
| v5p | 4.59e14 | $4.2 | 3.9e17 |
| v5e | 1.97e14 | $1.2 | 5.8e17 |
| TPU 型号 | bfloat16 FLOPs/s | Google Cloud 美元 / 小时 | FLOPs / $ |
|---|---|---|---|
| H100 | 9.9e14 | $10.8 | 3.3e17 |
| v5p | 4.59e14 | $4.2 | 3.9e17 |
| v5e | 1.97e14 | $1.2 | 5.8e17 |
Each TPU v5e has 16GB of HBM which will require us to shard our model fairly aggressively. Let’s start by thinking about some basic quantities that might matter for us:
每颗 TPU v5e 只有 16GB HBM,这要求我们相当激进地分片模型。先从一些可能对我们重要的基本量开始:
Question: How large are LLaMA 3-70B’s KV caches per token? You can assume we store them in int8. This determines how large our batch size can be on a given topology.
问题: LLaMA 3-70B 每个 token 的 KV cache 有多大?可以假设用 int8 存储。这决定了在给定拓扑下我们能开多大的 batch size。
Click here once you've thought it through!
想清楚之后点击这里!
LLaMA 3-70B has 8 KV heads, so the size per token is 2 * K * H * L = 2 * 8 * 128 * 80 = 160kB.
LLaMA 3-70B 有 8 个 KV 头,所以每个 token 的大小是 2 * K * H * L = 2 * 8 * 128 * 80 = 160kB。
Note just how big this is! If we have a sequence length of 32k tokens (as is common), this uses 160e3 * 32,768 = 5.3GB / sequence. For BS=240, this is 1.3TB! Since TPU v5e only have 16GB a piece, we would need about (70e9 + 1.3e12) / 16e9 = 86 TPU v5e chips to even fit this much memory. Also note how large this is compared to the 70GB of model parameters.
注意这有多大! 如果序列长度是 32k 个 token(很常见),这会占用 160e3 * 32,768 = 5.3GB / 序列。对 BS=240,就是 1.3TB!由于每颗 TPU v5e 只有 16GB,光是装下这么多显存就需要大约 (70e9 + 1.3e12) / 16e9 = 86 颗 TPU v5e。也注意一下,相比 70GB 的模型参数,这是多么大。
Question: Let’s say we want to serve L3 70B at batch size 32 and 8192 sequence length with everything (params and KVs) in int8. How much total memory will this use? What’s the smallest slice we could serve this on?
问题: 假设我们要以 batch size 32、序列长度 8192 服务 L3 70B,所有东西(参数和 KV)都用 int8。总共会用多少显存?能服务它的最小 slice 是多大?
Answer
答案
Since our KVs are 160e3 bytes in int8, our total KV memory is 160e3 * 8192 * 32 = 41.9e9 bytes. Our parameters are 70e9 bytes, since we have 1 byte per parameter. Thus, our total memory usage is 41.9e9 + 70e9 = 112GB.
在 int8 下我们的 KV 是 160e3 字节,所以 KV 总显存是 160e3 * 8192 * 32 = 41.9e9 字节。参数是 70e9 字节,因为每个参数 1 字节。因此总显存占用是 41.9e9 + 70e9 = 112GB。
The smallest slice we could use would have 112e9 / 16e9 = 7 TPUs, or (rounding to an even size), TPU v5e 4x2. This will be a tight fit and we might not be able to quite fit this accounting for other overhead, so we might need a 4x4 at minimum (or to drop the batch size).
能用的最小 slice 需要 112e9 / 16e9 = 7 颗 TPU,取整到偶数规模就是 TPU v5e 4x2。这会很紧,考虑到其他开销我们可能装不下,所以最少可能需要 4x4(或者降低 batch size)。
Question: At this batch size and quantization on a TPU v5e 4x2, roughly what latency would we expect per decode step? What throughput (tokens / sec / chip). What about a 4x4? Assume we perform our FLOPs in bfloat16 and everything is fully sharded.
问题: 在 TPU v5e 4x2 上、用这个 batch size 和量化,每个 decode 步的延迟大约是多少?吞吐是多少(tokens / 秒 / 芯片)?换成 4x4 呢?假设 FLOPs 用 bfloat16 做,且一切都完全分片。
Answer
答案
We can invoke the formula from the previous section that
我们可以套用上一节的公式:
$$\begin{align*} \tiny \text{Theoretical Step Time (General)} = \underbrace{\frac{\text{Batch Size} \times \text{KV Cache Size}}{\tiny \text{Total Memory Bandwidth}}}_{\text{Attention (always bandwidth-bound)}} + \underbrace{\max\left(\frac{2 \times \text{Batch Size} \times \text{Parameter Count}}{\text{Total FLOPs/s}}, \frac{\text{Parameter Size}}{\text{Total Memory Bandwidth}}\right)}_{\tiny \text{MLP (can be compute-bound)}} \end{align*}$$
Here our critical batch size will be about 120 since our parameters are in int8 but our FLOPs are in bfloat16. We could also manually calculate the RHS maximum, but that’s basically a calculation we’ve already done several times. So we’re well into the memory-bound regime for both our matmul and our FLOPs.
这里的临界 batch size 大约是 120,因为参数是 int8 而 FLOPs 是 bfloat16。我们也可以手算右边那个最大值,但那基本是我们已经算过好几遍的东西了。所以我们的矩阵乘法和 FLOPs 都稳稳处于显存受限区间。
Strictly looking at memory bandwidth then, our step time is basically (KV size + param size) / (8 * HBM bandwidth) = 112e9 / (8 * 8.2e11) = 17ms. So theoretically our step time is about 17ms. Our throughput would be 32 / .017 = 1882 tokens / sec, or 1882 / 8 = 235 tokens / sec / chip.
只看显存带宽的话,我们的步时基本是 (KV 大小 + 参数大小) / (8 * HBM 带宽) = 112e9 / (8 * 8.2e11) = 17ms。所以理论上步时约为 17ms。 吞吐是 32 / .017 = 1882 tokens / 秒,即 1882 / 8 = 235 tokens / 秒 / 芯片。
There’s one caveat here which is to check if we might be ICI bound on our matmuls. We could dedicate 2 axes to it here, so we’re ICI bound in theory when $Y > 2 * F / 2200 = 2 * 28672 / 2200 = 26$, so we’re golden!
这里有一个需要注意的点:检查我们的矩阵乘法是否会 ICI 受限。这里我们能把 2 条轴分给它,所以理论上当 $Y > 2 * F / 2200 = 2 * 28672 / 2200 = 26$ 时才 ICI 受限,我们完全没问题!
If we were to run on a 4x4, we’d still be fine ICI-wise, so our latency would drop to 17 / 2 = 8.5ms, but our throughput per-chip would remain the same.
如果跑在 4x4 上,ICI 方面仍然没问题,于是延迟降到 17 / 2 = 8.5ms,但每芯片吞吐保持不变。
Thinking about throughput
想一想吞吐
Let’s spend a little time thinking purely about throughput. When we optimize for throughput, we want to be compute bound, meaning we come close to utilizing all the TPU MXU capacity. Typically that means we want the batch size to be as large as possible, so we are doing as much work as possible.
我们花点时间纯粹从吞吐的角度想想。当我们为吞吐做优化时,我们希望是计算受限,也就是接近把 TPU 的全部 MXU 算力用满。通常这意味着我们要让 batch size 尽可能大,从而做尽可能多的工作。
Question: On TPU v5e, using bfloat16 weights and activations, how large do our batch sizes need to be for us to be compute-bound in our matmuls? What if we do int8 weights but perform our FLOPs in bfloat16? What about int8 weights with int8 FLOPs?
问题: 在 TPU v5e 上,用 bfloat16 权重和激活时,batch size 要多大我们的矩阵乘法才是计算受限?如果用 int8 权重但 FLOPs 用 bfloat16 呢?int8 权重配 int8 FLOPs 呢?
Answer
答案
As discussed in Section 7, for any bfloat16 matmul for which $B \ll D, F$ we have
如第 7 节所述,对任何满足 $B \ll D, F$ 的 bfloat16 矩阵乘法,我们有
$$\begin{equation*} T_\text{math} > T_\text{comms} \leftrightarrow \frac{2BDF}{2DF} \geq \frac{\text{TPU bfloat16 FLOPs/s}}{\text{HBM bandwidth}} = 240 \end{equation*}$$
When our weights are in int8, we lose a factor of 2 in the denominator, so we have $2BDF / DF = 2B > 240$, or equally $B > 120$, half the critical batch size from before. That’s really helpful for us! When we do int8 weights and int8 FLOPs, we have to use the int8 value for TPU FLOPs/s, which goes from 1.97e14 for bfloat16 to 3.94e14, nearly double. That means we’re back where we started at about $B > 240$.
当权重是 int8 时,分母丢了一个 2 倍因子,于是有 $2BDF / DF = 2B > 240$,也就是 $B > 120$,是之前临界 batch size 的一半。这对我们非常有帮助!当权重和 FLOPs 都是 int8 时,我们必须用 TPU int8 的 FLOPs/s,它从 bfloat16 的 1.97e14 变成 3.94e14,几乎翻倍。这意味着我们又回到起点,约 $B > 240$。
The case of int8 weights and bfloat16 FLOPs is quite common, since quantizing parameters losslessly is often easier than doing low-precision arithmetic.
int8 权重配 bfloat16 FLOPs 这种组合相当常见,因为无损地量化参数往往比做低精度算术更容易。
Question: What is the smallest TPU v5e topology we could serve LLaMA 3-70B on using bfloat16, int8, and int4 (both KVs and parameters) with 8k context? You can think of KV caches as negligibly small for this one.
问题: 在 8k 上下文下,分别用 bfloat16、int8 和 int4(KV 和参数都用各自精度)服务 LLaMA 3-70B,所需的最小 TPU v5e 拓扑是多大?这题里可以把 KV cache 视作小到可忽略。
Answer
答案
This is easy! If we’re OK with a tiny batch size then the only limit is fitting parameter memory in HBM, i.e. it is just ceil(num_params * sizeof(dtype) / HBM per TPU), or ceil(70e9 * sizeof(dtype) / 16e9) rounded to the nearest reasonable topology (some multiple of 2):
很简单!如果我们能接受很小的 batch size,那么唯一的限制就是把参数装进 HBM,也就是 ceil(num_params * sizeof(dtype) / HBM per TPU),即 ceil(70e9 * sizeof(dtype) / 16e9),再取整到最近的合理拓扑(2 的若干倍):
| dtype | param size | KV size / token (bytes) | min TPU v5es | actual min slice | remaining HBM for KV caches | num KV caches @ 8k |
|---|---|---|---|---|---|---|
| bf16 | 140GB | 324kB | 8.75 | 4x4 = 16 chips | 116 | 43 |
| int8 | 70GB | 162kB | 4.38 | 4x2 = 8 chips | 58 | 43 |
| int4 | 35GB | 81kB | 2.81 | 2x2 = 4 chips | 29 | 43 |
| dtype | 参数大小 | 每 token KV 大小(字节) | 最少 TPU v5e 数 | 实际最小 slice | 剩余给 KV cache 的 HBM | @8k 的 KV cache 份数 |
|---|---|---|---|---|---|---|
| bf16 | 140GB | 324kB | 8.75 | 4x4 = 16 颗 | 116 | 43 |
| int8 | 70GB | 162kB | 4.38 | 4x2 = 8 颗 | 58 | 43 |
| int4 | 35GB | 81kB | 2.81 | 2x2 = 4 颗 | 29 | 43 |
That’s pretty cool! It tells us we could fit LLaMA 70B on a TPU v5e 2x2 if we wanted to. Except you’ll notice the number of KV caches is very small. That’s our batch size! That means we’ll be getting terrible FLOPs utilization. We’d be very happy to use a larger topology in order to push our batch size up to 240.
相当酷!这说明只要愿意,我们甚至能把 LLaMA 70B 塞进一个 TPU v5e 2x2。不过你会注意到 KV cache 的份数非常少。那就是我们的 batch size!这意味着 FLOPs 利用率会非常差。我们会非常乐意用更大的拓扑把 batch size 推到 240。
Question: Assume we use the largest batch size that fits on these topologies, what latency could we expect for each generate step?
问题: 假设我们在这些拓扑上使用能装下的最大 batch size,每个 generate 步的延迟大约是多少?
Answer
答案
This is also easy, since we’re picking our batch size to fill up all our HBM! This is just a question of how long it takes to load a full TPU v5e’s worth of bytes into the MXU. This is just v5e HBM / v5e HBM memory bandwidth = 16GB / 8.2e11 = 19ms, so this is 19ms / step. Assuming our generations have a median length of 512 tokens, that is about 9s for each decode. Note that we could get marginally better latency with a smaller batch size, for instance if we only looked at model parameters in int4 our minimum latency is about 10ms / step, since HBM is no longer full.
这也不难,因为我们选的 batch size 会填满全部 HBM!问题就只是把一整个 TPU v5e 的字节量加载进 MXU 需要多久,也就是 v5e HBM / v5e HBM 显存带宽 = 16GB / 8.2e11 = 19ms,所以是 19ms / 步。假设生成的典型长度是 512 个 token,每次 decode 大约 9 秒。注意,用更小的 batch size 可以得到略好的延迟,例如只看 int4 的模型参数时,由于 HBM 不再被填满,最小延迟大约是 10ms / 步。
Takeaway: we can always lower bound decode latency by asking how long it takes to load all the model’s parameters from HBM into the MXU. When our KV caches are small, you can think about each layer as just loading the weights chunk-by-chunk and then discarding them. Unless we’re using large batch sizes or lots of inter-device comms, this is often a reasonable bound (within 1.5x). When our batch size is bigger, we need to model the KV cache loading as well, since that dominates the parameters.
结论: 我们总能用「把模型全部参数从 HBM 加载进 MXU 需要多久」来给出 decode 延迟的下界。当 KV cache 很小时,你可以把每一层想成只是逐块加载权重、用完就丢。除非用大 batch size 或大量设备间通信,这往往是个合理的界(误差在 1.5 倍以内)。当 batch size 更大时,我们还要把 KV cache 的加载也建模进去,因为它会超过参数。
Likewise, in the FLOPs-bound regime (e.g. training or big-batch inference), we can use the $$\text{Total FLOPs} / (N \cdot C) = 2 \cdot \text{param count} \cdot B / (N \cdot C)$$ lower bound, which assumes no communication.
同样地,在 FLOPs 受限的区间(例如训练或大 batch 推理),我们可以用 $$\text{Total FLOPs} / (N \cdot C) = 2 \cdot \text{param count} \cdot B / (N \cdot C)$$ 这个下界,它假设没有通信。
Question: For each of these, what throughput per chip does this give us (in terms of queries / chip)? You can assume our median decode length is 512 tokens.
问题: 对每一种情况,这给我们带来多少每芯片吞吐(以 queries / 芯片计)?可以假设典型的 decode 长度是 512 个 token。
Answer
答案
This is an important question because it’s exactly correlated with cost / token.
这个问题很重要,因为它与 cost / token 精确相关。
With our assumption about median decode length, our throughput is just $$B / (\text{per-step latency} \cdot \text{median steps} \cdot N) \approx 43 / (0.019 * 512 * N)$$. This gives us roughly $$(4.42 / N)$$ QPS, so plugging in $$N$$ we get:
在典型 decode 长度的假设下,我们的吞吐就是 $$B / (\text{每步延迟} \cdot \text{步数中位数} \cdot N) \approx 43 / (0.019 * 512 * N)$$。这大致给出 $$(4.42 / N)$$ QPS,代入 $$N$$ 得到:
| dtype | QPS / chip |
|---|---|
| bfloat16 | 0.27 |
| int8 | 0.55 |
| int4 | 1.11 |
| dtype | 每芯片 QPS |
|---|---|
| bfloat16 | 0.27 |
| int8 | 0.55 |
| int4 | 1.11 |
Note that this is rather optimistic since it totally ignores the working memory of the forward pass (memory allocated to activations and attention). This is not ridiculous with Flash Attention, but it is also not realistic. The real numbers are likely around 1/2 of this. For absolutely maximum throughput we would probably want to more than double the number of chips and increase the batch size significantly as well.
注意这相当乐观,因为它完全忽略了前向传播的工作显存(分配给激活和注意力的显存)。有 Flash Attention 时这不算离谱,但也不现实。真实数字很可能是它的一半左右。要追求绝对最大吞吐,我们大概需要把芯片数翻倍还多,并显著增大 batch size。
Question: How would our peak throughput change if we doubled our topology for each of the above examples?
问题: 如果把上面每个例子的拓扑都翻倍,我们的峰值吞吐会怎么变?
Answer
答案
If we used a 4x8 slice in bfloat16, we would have 372GB remaining for KV caches, which would let us increase our batch size to 140. Then since our step time would remain the same, we would have a throughput of 14.39 / num_chips, or
如果用 bfloat16 跑在 4x8 slice 上,我们会剩 372GB 给 KV cache,能让我们把 batch size 提到 140。由于步时不变,吞吐会是 14.39 / num_chips,即
| dtype | QPS / chip |
|---|---|
| bfloat16 (on 4x8) | 0.44 |
| int8 (on 4x4) | 0.90 |
| int4 (on 2x4) | 1.80 |
| dtype | 每芯片 QPS |
|---|---|
| bfloat16 (on 4x8) | 0.44 |
| int8 (on 4x4) | 0.90 |
| int4 (on 2x4) | 1.80 |
A further increase would give an even bigger win! The big takeaway is that the smallest topology is not the most performant topology in all cases, if we’re limited by KV cache size.
再往上加还能带来更大的收益!核心结论是:如果被 KV cache 大小限制,最小的拓扑并不总是性能最好的拓扑。
Question: Now let’s dig into the question of sharding. Let’s say we wanted to serve in bfloat16 on a TPU v5e 4x8. What sharding would we use for our model on a TPU v5e 4x8 during generation? Can we avoid being communication bound?
问题: 现在深入看一下分片问题。假设我们要在 TPU v5e 4x8 上用 bfloat16 服务。在 TPU v5e 4x8 上做 generation 时,我们会用哪种分片?我们能避免通信受限吗?
Answer
答案
As discussed in the previous section, we only really have one option for sharding during generation: model parallelism. How much can we do before we become communication bound? As we’ve discussed in the previous section, our models become communication bound roughly when
如上一节所述,generation 时分片其实只有一种选择:模型并行。在变成通信受限之前我们最多能做多少?如上一节所讨论的,当
$$Y > \frac{F \cdot M_Y}{2200}$$
For LLaMA 3-70B we have F = 28,672, so if we do 2 axes of model sharding this gives us roughly $$Y = 28672 \cdot 2 / 2200 = 26$$, so in general we could scale up to about 16 chips without being communication bound, which lets us use a 4x4 but not a 4x8. Generally, since we do not perfectly overlap computation, even this estimate is overly optimistic.
对 LLaMA 3-70B,F = 28,672,所以如果做 2 条轴的模型分片,大约是 $$Y = 28672 \cdot 2 / 2200 = 26$$,即一般可以扩展到约 16 颗芯片而不受通信限制,这让我们能用 4x4 但用不了 4x8。一般来说,由于我们的计算重叠并不完美,即便这个估计也过于乐观了。
Takeaway: we cannot actually serve on a 4x8 with pure model parallelism. The best we can do here is a 4x2 or maybe a 4x4.
结论:纯模型并行其实无法在 4x8 上服务。 这里我们能做的最好是 4x2,或者也许 4x4。
However, as we’ve discussed, when our batch size is small we can often do more model parallelism without significantly hurting throughput, since our model is memory-bandwidth-bound and not FLOPs bound. We said before that this value is roughly $Y=F / (8\cdot B)$, so if we did batch size 64, we could in theory go up to Y = 28,672 / (8 * 64) = 56 way model parallelism before we become ICI-bound. To sanity check this, we can look at $T_\text{ici comms}$, $T_\text{hbm comms}$, and $T_\text{math}$ for a single matmul. We clearly have:
不过,如我们之前所讨论的,当 batch size 较小时,我们往往能在不显著伤害吞吐的前提下做更多模型并行,因为模型是显存带宽受限而不是 FLOPs 受限。我们之前说过这个值大约是 $Y=F / (8\cdot B)$,所以如果 batch size 取 64,理论上最多可以做到 Y = 28,672 / (8 * 64) = 56 路模型并行才 ICI 受限。为验证这一点,我们可以看单次矩阵乘法的 $T_\text{ici comms}$、$T_\text{hbm comms}$ 和 $T_\text{math}$。显然有:
$$\begin{align*}T_\text{ici comms} = \frac{2BD}{W_\text{ici}} && T_\text{hbm comms} = \frac{2DF}{Y \cdot W_\text{hbm}} && T_\text{math} = \frac{2BDF}{Y \cdot C}\end{align*}$$
For a 4x8, this would give us $T_\text{ici comms}$ = (2 * 64 * 8192) / 9e10 = 11us, $T_\text{hbm comms}$ = (2 * 8192 * 28,672) / (32 * 8.2e11) = 18us, and $T_\text{math}$ = (2 * 64 * 8192 * 28,672) / (32 * 1.97e14) = 4us, so in theory we’re still HBM bandwidth bound, which is great! Note that scaling up from a 4x4 to a 4x8 probably isn’t helpful from a throughput standpoint, but it’ll reduce our latency!
对 4x8 来说,这给出 $T_\text{ici comms}$ = (2 * 64 * 8192) / 9e10 = 11us、$T_\text{hbm comms}$ = (2 * 8192 * 28,672) / (32 * 8.2e11) = 18us、$T_\text{math}$ = (2 * 64 * 8192 * 28,672) / (32 * 1.97e14) = 4us,所以理论上我们仍是 HBM 带宽受限,这很好!注意,从 4x4 扩到 4x8 从吞吐角度大概没什么帮助,但会降低延迟!
If we look at the int8 and int4 configs, we can do those with pure model parallelism. So we’ve hit a point at which quantization actually gives us a meaningful advantage beyond faster FLOPs: it lets us use a larger batch size before we become comms-bound. So the end of this story is that we can’t achieve peak throughput on a 4x8, but for the int8 and int4 configs we could do pure model parallelism.
如果看 int8 和 int4 的配置,我们可以用纯模型并行做到。所以到了这个地步,量化除了更快的 FLOPs 之外,还真的给了我们有意义的优势:它让我们在变成通信受限之前能用更大的 batch size。所以这个故事的结局是:在 4x8 上我们无法达到峰值吞吐,但对 int8 和 int4 配置,我们可以用纯模型并行。
Tip: the maximum amount of useful model parallelism depends on $$d_{ff}$$ and the number of axes over which you’re sharding your model. The maximum value usually ranges between 8 and 32 depending on the model size. You can scale beyond this limit to improve latency at some throughput cost.
提示:有用的模型并行的上限取决于 $$d_{ff}$$ 以及你分片模型所跨的轴数。这个最大值通常随模型大小在 8 到 32 之间。你可以越过这个上限来以一定的吞吐代价改善延迟。
What about prefill?
那 prefill 呢?
We’ve mostly ignored prefill here because it’s much simpler. Let’s put a couple of concepts together and think about the end-to-end picture.
这里我们基本忽略了 prefill,因为它简单得多。我们把几个概念放一起,想想端到端的图景。
Question: Assume we achieve a 40% FLOPs utilization during prefill. How long will a prefill of length 8192 take on 16 TPU v5e chips?
问题: 假设 prefill 时我们达到 40% 的 FLOPs 利用率。在 16 颗 TPU v5e 上,长度 8192 的 prefill 要花多久?
Answer
答案
At 8k tokens, we are solidly compute bound, so we just need to reason about FLOPs. We know our model has 70e9 parameters so each forward pass uses 2 * 70e9 * B FLOPs. Assuming 40% MFU (FLOPs utilization), this gives us a runtime of about 2 * 70e9 * 8192 / (16 * 1.97e14 * 0.4) = 0.91s. Compared to the numbers we’ve been looking at before, that’s actually quite a lot!
在 8k 个 token 时,我们稳稳是计算受限,所以只需考虑 FLOPs。我们知道模型有 70e9 个参数,所以每次前向传播用 2 * 70e9 * B 次 FLOPs。假设 40% MFU(FLOPs 利用率),运行时间大约是 2 * 70e9 * 8192 / (16 * 1.97e14 * 0.4) = 0.91s。和我们之前看的那些数字相比,这其实相当多了!
Question: Assume we have a median prefill length of 8192 tokens and a median decode length of 4096 tokens. Say we have a generate batch size of 32. On average how many sequences finish decoding per step? On average how many tokens are evicted from our KV cache each step?
问题: 假设 prefill 长度中位数是 8192 个 token,decode 长度中位数是 4096 个 token。假设 generate batch size 是 32。平均每一步有多少条序列完成解码?平均每一步有多少 token 从 KV cache 中被淘汰?
Answer
答案
This is kind of straightforward. Since we have a median decode length of 4096 tokens, a sequence will finish roughly every 1 / 4096 tokens. Given a batch size of 32, this means we have 32 / 4096 sequences evicted per step. Since our KV cache length is roughly 8192 + 4096, this is 32 * (8192 + 4096) / 4096 = 96 tokens evicted per step. The general formula is $B * (P + G) / G$ where $P$ and $G$ are the prefill and generate lengths.
这比较直接。由于 decode 长度中位数是 4096 个 token,平均每 1 / 4096 个 token 就有一条序列完成。给定 batch size 32,这意味着每步有 32 / 4096 条序列被淘汰。由于 KV cache 长度大约是 8192 + 4096,所以每步淘汰 32 * (8192 + 4096) / 4096 = 96 个 token。一般公式是 $B * (P + G) / G$,其中 $P$ 和 $G$ 分别是 prefill 和 generate 的长度。
Question: Assume we do disaggregated serving with a median prefill length of 8192 and a median decode length of 512. Assume the prefill and generate latencies calculated above in bfloat16. What ratio of prefill servers will you need to keep both fully saturated.
问题: 假设我们做分离式服务,prefill 长度中位数是 8192,decode 长度中位数是 512。假设上面算的 prefill 和 generate 延迟都按 bfloat16。要让两边都保持饱和,需要多少 prefill 服务器比例?
Answer
答案
This is kind of a fun question. Let $P$ be the number of prefill servers and $G$ be the number of generate servers. So generally speaking, this is a pipeline problem where we feed sequences in at a rate of P / prefill_latency and consume them at a rate of B * G / (generate_latency * median_decode_length). We had calculated 910ms per prefill step and 19ms per decode step at batch size 43 (let’s call that 32). Therefore we need P / 0.91 = 32 * G / (0.019 * 512) or P = 3G, i.e. we need about 3 times more prefill servers than generation servers!
这题挺有意思。设 $P$ 为 prefill 服务器数量,$G$ 为 generate 服务器数量。一般来说这是个流水线问题:我们以 P / prefill_latency 的速率喂入序列,以 B * G / (generate_latency * median_decode_length) 的速率消费它们。我们之前算出每步 prefill 是 910ms、batch size 43(就当它是 32)下每步 decode 是 19ms。因此需要 P / 0.91 = 32 * G / (0.019 * 512),即 P = 3G,也就是我们需要大约 3 倍于 generate 服务器的 prefill 服务器!
Visualizing the Latency Throughput Tradeoff
可视化延迟-吞吐取舍
Sticking with LLaMA 70B for a second, let’s actually look at the latency and throughput for different batch sizes during generation. As we showed in the previous section for PaLM models, this gives us a Pareto frontier for throughput/latency. Let’s assume 16-way tensor parallelism since that’s a reasonable bound on what we can use while staying compute-bound in the MLP blocks. We’ll use a TPU v5e 4x4 topology here. The slider controls the sequence length so you can see the effect of larger KV caches.
继续用 LLaMA 70B,我们实际看看 generation 时不同 batch size 下的延迟和吞吐。如上一节对 PaLM 模型所展示的,这给出了吞吐 / 延迟的 Pareto 前沿。假设我们用 16 路张量并行,因为这是我们在 MLP 块保持计算受限的前提下能用的一个合理上限。这里我们用 TPU v5e 4x4 拓扑。滑块控制序列长度,让你看到更大 KV cache 的影响。
- See how dramatic the tradeoff is between cost and latency. At the cost of doubling per-token latency, we can achieve a roughly 100x reduction in per-token cost. Also, our latency can range anywhere from 5.5ms with low batch size to 20 ms with very large batches.
- Note how at 2k context the throughput effectively plateaus at around 1 token / ms / chip when it hits the BS 120 roofline (120 here because we do int8 weights but bf16 FLOPs). As the sequence length increases, however, we can no longer fit this batch size in memory, so we never hit the point of full saturation.
- Note how much higher the latency is at large batch sizes for the same throughput, since KV loading becomes dominant (instead of parameter loading).
- 看看成本与延迟之间的取舍有多剧烈。 以每 token 延迟翻倍为代价,我们可以把每 token 成本降低大约 100 倍。而且,我们的延迟可以从低 batch 时的 5.5ms 一直到超大 batch 时的 20ms。
- 注意在 2k 上下文时,吞吐在撞上 BS 120 roofline 后基本在 1 token / ms / 芯片 附近走平(这里是 120,因为我们用 int8 权重但 bf16 FLOPs)。然而随着序列长度增加,我们再也装不下这个 batch size,所以永远达不到完全饱和的点。
- 注意在相同吞吐下,大 batch size 的延迟要高得多,因为 KV 加载(而不是参数加载)成为主导。
We can understand this better by breaking down the sources of cost and latency into param loading time, KV loading time, and FLOPs time. The shaded region is where we expect to be compute-bound in our MLP blocks.
我们可以把成本和延迟来源拆成参数加载时间、KV 加载时间和 FLOPs 时间,从而更好地理解这一点。阴影区域是我们预期 MLP 块计算受限的地方。
This tells quite a story. You can see that initially, parameter loading represents the vast majority of the latency, until the batch size becomes large enough that FLOPs and KV loading become more significant. Notably, at all sequence lengths greater than 2048, we spend more time on KV cache loading than we do on FLOPs! So while we can improve our hardware utilization by increasing batch size, at long context lengths KV loading always dominates the total step time.
这里讲了一个很完整的故事。你可以看到,起初参数加载占延迟的绝大部分,直到 batch size 大到 FLOPs 和 KV 加载也变得重要。值得注意的是,在所有大于 2048 的序列长度下,我们花在 KV cache 加载上的时间都比 FLOPs 更多!所以尽管增大 batch size 可以提升硬件利用率,但在长上下文下,KV 加载总是主导总步时。
Takeaway: for LLaMA 3-70B, we are strongly KV cache memory bandwidth-bound (and HBM-bound) in almost all of these configurations, highlighting just how important reducing KV cache size is for generation throughput. Also note just how dramatic the latency/throughput tradeoff remains here.
结论: 对 LLaMA 3-70B,在几乎所有配置下我们都强烈受 KV cache 显存带宽限制(也是 HBM 受限),这凸显了减小 KV cache 大小对 generation 吞吐有多么重要。也注意一下这里的延迟 / 吞吐取舍依然多么剧烈。
The code for this is quite simple.
相关代码相当简单。
Here’s the code for computing these rooflines:
下面是计算这些 roofline 的代码:
import numpy as np
num_chips = 16 # we fix 16 as the amount of total model parallelism we do
bytes_per_param = 1 # int8 means 1 byte per param
param_count = 70e9
param_size = bytes_per_param * param_count
sequence_length = 8192 # can vary this
hbm_bandwidth = 8.20E+11 # v5e
flops = 1.97E+14 # v5e
def kv_cache_size(bs):
return 2 * bs * 128 * 8 * 80
def min_topology(bytes):
return 2 ** np.ceil(np.log2(bytes / 16e9))
def get_max_batch_size(
num_chips: int,
sequence_length: int,
param_size: float,
) -> int:
batch_sizes = np.arange(1, 1024, 4)
kv_sizes = kv_cache_size(sequence_length * batch_sizes)
required_chips = min_topology(kv_sizes + param_size)
max_idx = np.where(required_chips <= num_chips)[0][-1]
return max_idx
max_idx = get_max_batch_size(
num_chips=num_chips,
sequence_length=sequence_length,
param_size=param_size,
) # get the largest batch size that can fit
batch_sizes = np.arange(1, 512, 1)[:max_idx]
kv_sizes = kv_cache_size(sequence_length * batch_sizes)
kv_comms_time = kv_sizes / (num_chips * hbm_bandwidth)
param_comms_time = param_size / (num_chips * hbm_bandwidth)
param_comms_time = np.asarray([param_comms_time] * batch_sizes.shape[0])
flops_time = 2 * param_size * batch_sizes / (num_chips * flops) # roughly true in a 2ND sense
mlp_time = np.maximum(flops_time, param_comms_time)
attn_time = kv_comms_time # always bandwidth-bound for generate
latency = 1000 * (mlp_time + attn_time)
throughput = batch_sizes / (latency * num_chips)Note how we very explicitly break out latency into two sources: KV loading and param loading, and how the latency is either bound by FLOPs or comms, whichever is bigger.
注意我们非常显式地把延迟拆成两个来源:KV 加载和参数加载;并且延迟要么被 FLOPs 约束、要么被通信约束,取两者中更大的那个。
Worked Problems
实战习题
Here are a few worked problems. Some of these repeat things that are worked above, but might be pedagogically useful.
下面是几道实战题。有些是重复上面做过的内容,但可能对教学有帮助。
Question 1: How many FLOPs does each forward pass for LLaMA 3-405B use per-token? Assuming we’re FLOPs bound, what is a lower bound on a single forward pass on N chips on TPU v5e? What if we’re comms bound? Ignore the fact that the model does not fit on a single chip.
问题 1: LLaMA 3-405B 每次前向传播每 token 用多少 FLOPs?假设 FLOPs 受限,在 TPU v5e 上 N 颗芯片做一次前向传播的下界是多少?如果通信受限呢?忽略模型装不进单颗芯片这一事实。
Question 2: Assume we want to serve LLaMA 3-8B with BS240 using int8 weights and int8 KV caches. How many bytes are used by (a) model parameters (b) KV caches and (c) peak working activations (roughly)? What’s the smallest topology we can run this on?
问题 2: 假设我们要用 BS240、int8 权重和 int8 KV cache 服务 LLaMA 3-8B。(a) 模型参数、(b) KV cache、(c) 峰值工作激活(粗略)各占用多少字节?能跑它的最小拓扑是多大?
Question 3: How would you serve LLaMA 3-405B on TPU v5e? Assume int8 weights and bfloat16 FLOPs. Let’s say we have a firm limit of 15ms / token, what’s the highest throughput configuration we could achieve? What is the theoretical minimum step time?
问题 3: 你会如何在 TPU v5e 上服务 LLaMA 3-405B?假设 int8 权重、bfloat16 FLOPs。假设我们有 15ms / token 的硬性上限,我们能达到的最高吞吐配置是什么?理论最小步时是多少?
Question 4: The best way to learn about LLMs is to implement one from scratch. Building the full training pipeline is annoying and expensive, but turning trained weights you can download from HuggingFace into a working inference implementation is extremely instructive. Roughly speaking, you should try to do the following:
问题 4: 学习 LLM 最好的方式就是从零实现一个。搭完整训练管线又烦又贵,但把从 HuggingFace 下载的训练好权重变成一个能用的推理实现,极具教育意义。大致来说,你应该尝试做以下几件事:
- Download LLaMA 3 8B weights and load them in Colab. Visualize the weights (I like TreeScope for this) and see if you can identify each tensor. Count the # of parameters. Does it match what you expect? How many are in the MLP vs. attention?
- 下载 LLaMA 3 8B 权重并在 Colab 里加载。可视化这些权重(我喜欢用 TreeScope),看看能不能认出每个张量。数一数参数量,和你预期的一致吗?MLP 和注意力各占多少?
- Implement the full forward pass of the model. You don’t need to worry about prefill/decode or anything fancy. Just get to the point where you can feed in a sequence and get plausible next-token probabilities out. You should be able to feed a prompt in, get probabilities out, pick the highest probability one, and repeat. This is slow but functional sampling. Try to only use the LLaMA-3 paper as a reference, but the
jax-ml/jax-llm-examples/llama3repo can be a good reference for details (be careful of positional embeddings and causal masking). Your goal is to get coherent tokens out of the model. You can run this all on a single TPU if you can get more than 16GiB of HBM.
- 实现模型的完整前向传播。不需要操心 prefill/decode 之类花哨的东西,只要能喂进一段序列、得到看起来合理的 next-token 概率即可。你应该能喂进一个 prompt、得到概率、挑概率最高的那个,然后重复。这是慢但能用的采样。尽量只用 LLaMA-3 论文做参考,但
jax-ml/jax-llm-examples/llama3仓库在细节上是个不错的参考(小心位置编码和因果 mask)。你的目标是从模型里得到连贯的 token。如果你能搞到超过 16GiB 的 HBM,这一切都能在单颗 TPU 上跑。
- Now it’s time to make this faster. Implement KV caching, where you can save the key/value activations from a forward pass and attend to them in a subsequent one. This will make your sampling loop much faster.
- 现在让它更快。实现 KV caching,把一次前向传播的 key/value 激活保存下来,在后续前向传播中 attend 到它们。这会让你的采样循环快得多。
- Implement separate prefill and decode servers. You can do these on different sets of chips. Prefill handles just a single prompt at once, then sends it to a batched decode server that handles batches of tokens.
- 实现分离的 prefill 和 decode 服务器。可以把它们放在不同的芯片组上。prefill 一次只处理单个 prompt,然后把它发给处理批量 token 的 decode 服务器。
- Implement Flash Attention in Pallas. This will make attention much more efficient. I think it’s good to try and do this without a reference.
- 在 Pallas 里实现 Flash Attention。这会让注意力高效得多。我觉得不参考现成实现、自己试着做一遍是很好的练习。
讨论
用 GitHub 账号留言;评论保存在公开仓库chengshu-blog-discussions的 Discussions 里。也可通过 RSS 订阅后续文章。