如何扩展你的模型(6):在 TPU 上训练 LLaMA 3
中英对照第 6 篇:把前面的理论套到 LLaMA 3 上——需要多少 TPU、要训多久、大概花多少钱。左栏原文,右栏译文。
本篇属于系列 如何扩展你的模型(How To Scale Your Model) · 第 6 篇
原文: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)》的第 6 篇,共 13 篇。系列目录。
原文以 MIT 许可证发布,版权归 Google LLC;本译文仅作学习交流之用,如有错漏以原文为准。
Our goal in this section is to apply results from the previous section to a very practical problem: training the LLaMA 3 family (herd) of models. Unlike the previous sections we want you to do a lot of this work yourself. For this reason, we’ve hidden the answers to each section so you can try to answer it first. Try grabbing a pen and doing it by hand!
本节的目标是把前一节的结果应用到一个非常实际的问题上:训练 LLaMA 3 系列(herd)模型。与前面几节不同,我们希望你亲自动手完成大部分工作。因此,我们把每部分的答案都藏了起来,让你先自己试着回答。拿支笔,手算一下试试!
What does LLaMA 3 look like?
LLaMA 3 长什么样?
The LLaMA-3 model family[llama3] includes 3 main models: LLaMA 3 8B, 70B, and 405B. We’ll mostly focus on 70B, and leave 8B and 405B for you to explore in the problem section at the end. Here’s the architecture for LLaMA 3-70B, taken from the LLaMA HuggingFace page.
LLaMA-3 模型系列[llama3]包含 3 个主要模型:LLaMA 3 8B、70B 和 405B。我们主要关注 70B,把 8B 和 405B 留到最后的习题部分去探索。下面是 LLaMA 3-70B 的架构,取自 LLaMA 的 HuggingFace 页面。
| 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 |
To highlight how easy this is to find, here’s the config itself, along with a mapping:
为了说明这有多容易找到,下面给出 config 本身以及一份对照说明:
It’s useful to make a big table with these numbers for many different open-source LLMs, so you can quickly compare the design decisions they’ve made.
为许多不同的开源 LLM 整理一张包含这些数字的大表会很有用,这样你就能快速比较它们做出的设计决策。
Counting parameters and FLOPs
数参数与 FLOPs
Question: From this table, can we calculate the LLaMA 3-70B parameter count? 🤫 Let’s apply the content of Section 4 and see if we can get 70B!
问题: 根据这张表,我们能算出 LLaMA 3-70B 的参数量吗?🤫 让我们套用第 4 节的内容,看看能不能得出 70B!
| param | formula | count |
|---|---|---|
| FFW params | d_model * d_ff * 3 (for SwiGLU gate, up, and down projections) * n_layers | 8,192 * 8,192 * 3.5 * 3 * 80 = 56.3e9 |
| Vocab params | 2 (input and output embeddings) * n_embeddings * d_model | 2 * 128,256 * 8,192 = 2.1e9 |
| Attention params | n_layers * [ 2 (for q embedding and concatenated output projection) * d_model * n_heads * d_qkv + 2 (for k and v) * d_model * n_kv_heads * d_qkv] | 80 * (2 * 8,192 * 64 * 128 + 2 * 8,192 * 8 * 128) = 12e9 |
| 56.3e9 + 2.1e9 + 12e9 = 70.4e9 |
| 参数 | 公式 | 数量 |
|---|---|---|
| FFW 参数 | d_model * d_ff * 3(对应 SwiGLU 的 gate、up、down 三个投影)* n_layers | 8,192 * 8,192 * 3.5 * 3 * 80 = 56.3e9 |
| 词表参数 | 2(输入和输出 embedding)* n_embeddings * d_model | 2 * 128,256 * 8,192 = 2.1e9 |
| 注意力参数 | n_layers * [ 2(q 投影和拼接后的输出投影)* d_model * n_heads * d_qkv + 2(k 和 v)* d_model * n_kv_heads * d_qkv] | 80 * (2 * 8,192 * 64 * 128 + 2 * 8,192 * 8 * 128) = 12e9 |
| 56.3e9 + 2.1e9 + 12e9 = 70.4e9 |
That’s great! We get the number we expect. You’ll notice as expected that the FFW parameters totally dominate the overall parameter count, although attention is non-trivial.
很好!我们得到了预期的数字。如你所料,FFW 参数在总参数量中完全占主导,虽然注意力也不是个小数目。
Takeaway: The 3 big weight matrices in the MLP block are so much larger than all the other arrays in the Transformer that we can typically almost ignore all other parameters when reasoning about model memory or FLOPs. For LLaMA 3-70B, they represent 56B of 70B parameters.
结论: MLP 块里的 3 个大型权重矩阵比 Transformer 中其他所有数组都大得多,以至于在估算模型显存或 FLOPs 时,我们通常可以几乎忽略其他所有参数。对 LLaMA 3-70B 来说,它们在 70B 参数中占了 56B。
Let’s look at FLOPs now! Remember the general rules for training from Section 4.
现在我们来看 FLOPs!记得第 4 节里关于训练的一般规则。
Question: How many FLOPs does LLaMA-3 perform per token per training step? This helps us determine how expensive the whole training process will be.
问题: LLaMA-3 在每个训练步、每个 token 上要做多少次 FLOPs?这有助于我们判断整个训练过程有多贵。
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: As shown in Section 4, we do roughly $$6 \cdot \text{param count}$$ FLOPs per token, so here that’s roughly 6 * 70e9 = 4.2e11 FLOPs / token. That’s about half a TFLOP per token per step. Assuming we’re compute-bound, this should take roughly 4.2e11 / 4.59E+14 = 1ms on a single TPU v5p chip, assuming perfect FLOPs utilization.
答案: 如第 4 节所示,每个 token 大约做 $$6 \cdot \text{param count}$$ 次 FLOPs,所以这里大约是 6 * 70e9 = 4.2e11 FLOPs / token,也就是每步每 token 约半个 TFLOP。假设计算受限,在单颗 TPU v5p 芯片上(假定 FLOPs 利用率完美)大约需要 4.2e11 / 4.59E+14 = 1ms。
Question: LLaMA 3 was trained for about 15 trillion tokens. How many FLOPs is that total?
问题: LLaMA 3 训练了约 15 万亿个 token。总共有多少次 FLOPs?
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: That’s easy, it’s just 4.2e11 * 15e12 = 6.3e24 FLOPs total. 6.3 yottaFLOPs. That’s a lot! On a single TPU this would take 6.3e24 / 4.59E+14 = 435 years. That’s also a lot!
答案: 很简单,总共就是 4.2e11 * 15e12 = 6.3e24 FLOPs,即 6.3 yottaFLOPs。这非常多!在单颗 TPU 上要花 6.3e24 / 4.59E+14 = 435 年。这同样非常多!
Question: Let’s say we wanted to train on a full TPU v5p pod with 16x20x28 = 8960 chips. How long would this take to train at 40% MFU in bfloat16, assuming we are compute-bound?
问题: 假设我们要在一个 16x20x28 = 8960 颗芯片的完整 TPU v5p pod 上训练。假设计算受限、用 bfloat16、MFU 为 40%,训练要花多久?
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: We know that each TPU v5p can perform 4.59e14 FLOPs / second. At 40% MFU, this will take about T = 6.3e24 / (8960 * 4.59e14 * 0.4) = 3.8e6 seconds. This is about 44 days! That’s fairly reasonable, assuming we can actually achieve 40% MFU.
答案: 我们知道每颗 TPU v5p 每秒能做 4.59e14 次 FLOPs。在 40% MFU 下,大约需要 T = 6.3e24 / (8960 * 4.59e14 * 0.4) = 3.8e6 秒。这大约是 44 天! 假如真能达到 40% MFU,这还挺合理。
Question: LLaMA 3-70B was pretrained with a batch size of about 4M tokens. How many TPUs do we need at minimum to train with this batch size? You can assume bfloat16 parameters and float32 optimizer state, and that you checkpoint activations 2 times per layer (at the beginning of the layer and after attention).
问题: LLaMA 3-70B 预训练时 batch size 约为 4M token。要用这个 batch size 训练,我们最少需要多少颗 TPU?可以假设参数为 bfloat16、优化器状态为 float32,并且每层对激活做 2 次检查点(层开始处和注意力之后)。
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: This question is primarily asking about memory usage, since that’s the only strict constraint on available compute. During training, we have three primary uses of HBM: model parameters, optimizer state, and activation checkpoints. If we assume bfloat16 weights, float32 optimizer state, and a very conservative activation checkpointing scheme (2 times per layer), we have:
答案: 这道题主要问的是显存占用,因为那是对可用算力的唯一硬约束。训练时 HBM 主要有三种用途:模型参数、优化器状态和激活检查点。假设权重为 bfloat16、优化器状态为 float32,并采用一个非常保守的激活检查点方案(每层 2 次),我们有:
| Params | 2 * 70GB | ~140GB | | Optimizer State | 8 * 70GB | ~560GB | | Activation Checkpoints | 2 * 8192 * 4e6 * 2 * 80 | ~10.5TB | | Total | | ~11.2TB |
| 参数 | 2 * 70GB | ~140GB | | 优化器状态 | 8 * 70GB | ~560GB | | 激活检查点 | 2 * 8192 * 4e6 * 2 * 80 | ~10.5TB | | 合计 | | ~11.2TB |
The total here is about 11.2TB. You notice that activation checkpointing strongly dominates the memory picture, even with a very conservative checkpointing scheme. We could technically go to 1 checkpoint per layer, or do microbatching, but this is a reasonable picture. With these assumptions, since each TPU v5p has 96GB of HBM, we need 11.2e12 / 96e9 = 117 TPUs. That’s not very much actually!
这里合计约 11.2TB。你会注意到,即便用了非常保守的检查点方案,激活检查点依然在显存占用中占绝对主导。理论上我们可以降到每层 1 个检查点,或者用 microbatching,但这个图景是合理的。在这些假设下,由于每颗 TPU v5p 有 96GB HBM,我们需要 11.2e12 / 96e9 = 117 颗 TPU——其实并不多!
Why wouldn’t we do this? Well, because it would take us 44 days * 8960 / 117 = 3369 days to train. That’s nearly ten years. That’s a lot. Still, this makes it clear that we’re using these large clusters not because we’re bound by memory but rather because we need the extra FLOPs. Also, since we’re checkpointing so infrequently, we’re doing close to 8ND FLOPs instead of 6ND.
那我们为什么不这么做? 因为那样训练要花 44 days * 8960 / 117 = 3369 天,将近十年。这太久了。 不过这也清楚地说明:我们使用这些大集群,并不是因为受显存限制,而是因为我们额外需要那些 FLOPs。另外,由于检查点做得这么稀少,我们实际执行的接近 8ND 而不是 6ND 次 FLOPs。
Question: Under the same assumptions as the question above, if we use 8960 TPU v5p chips, how much memory will we use per-chip?
问题: 在与上一题相同的假设下,如果我们用 8960 颗 TPU v5p 芯片,每颗芯片会占用多少显存?
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: Our total memory is still about 11.2TB, so per-chip we’ll be using about 1.3GB per chip, which is basically nothing. If we did much more aggressive checkpointing, e.g. 12 checkpoints per layer, we’d still only be at 8GB per chip. We’re nowhere near being memory bound during training at these scales.
答案: 总显存仍然是约 11.2TB,所以平均每颗芯片约用 1.3GB,基本可以忽略。如果做激进得多的检查点(例如每层 12 个),每颗芯片也只有 8GB。在这个规模下训练,我们离显存受限还差得远。
Takeaways: It is technically possible to train even very large models on very small topologies, with the caveat that they will likely take a long time. Being able to calculate the total FLOPs of a training run allows us to ballpark its training time by assuming a modest MFU and a known topology.
要点: 从技术上讲,即使在很小的拓扑上也能训练非常大的模型,前提是它们大概率要跑很久。能够算出一次训练的总 FLOPs,让我们可以假设一个适中的 MFU 和已知拓扑,粗略估出训练时间。
How to shard LLaMA 3-70B for training
如何为训练分片 LLaMA 3-70B
Let’s stick to our setting from above and say we want to train LLaMA 3-70B with 4M token batch size (1024 sequences of length 4096 per batch) on a TPU v5p pod of 8960 chips. Let’s discuss what the best sharding strategy is for this model.
沿用上面的设定:我们要在 8960 颗芯片的 TPU v5p pod 上,用 4M token 的 batch size(每批 1024 条长度为 4096 的序列)训练 LLaMA 3-70B。我们来讨论这个模型的最佳分片策略。
Question: Under the assumptions above, can we train our model with FSDP alone? To start, let’s say we can’t do any sequence/context parallelism. This should be the first idea you have, since it’s simple and will introduce no extra communication if it works.
问题: 在上述假设下,我们能只用 FSDP 训练这个模型吗?先假设我们不能做任何序列 / 上下文并行。这应该是你第一个想到的方案,因为它简单,而且如果可行就不会引入额外通信。
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: This answer will be a little pedantic. As noted above, LLaMA 3-70B is initially trained with sequences of length 4K, so the batch size of 4M tokens gives us a sequence batch size of 1024. That means we can only really do pure data parallelism/FSDP up to 1024 chips because that’s how many sequences we have to do data parallelism over. So the answer in the simple sense of “fully data parallelism with no extra communication” is no. The next question will answer a slightly less pedantic version of this.
答案: 这个回答会有点抠字眼。如上所述,LLaMA 3-70B 最初是用长度 4K 的序列训练的,所以 4M token 的 batch size 给出的是 sequence batch size 1024。这意味着我们最多只能在 1024 颗芯片上做纯数据并行 / FSDP,因为只有这么多条序列可供数据并行。所以,从「不引入额外通信的纯数据并行」这个简单意义上说,答案是:不行。下一题会给出一个不那么抠字眼的版本。
Question: Let’s relax the requirement of not doing any sequence sharding. If we allow ourselves to do FSDP over both the batch and sequence axes, can we train LLaMA 3-70B with only FSDP on 8960 chips?
问题: 我们放宽「不做任何序列分片」的要求。如果我们允许自己在 batch 和 序列两条轴上都做 FSDP,那么能在 8960 颗芯片上只用 FSDP 训练 LLaMA 3-70B 吗?
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: Now that we’re allowing ourselves to do sequence/context parallelism as well, we can scale up way more. First let’s calculate our per-device batch size. If we do 8960-way FSDP, we end with a per-TPU batch size of 4 * 1024 * 1024 / 8960 = 468 tokens. We know from the previous section that we become ICI-bound by FSDP when $$\text{per device batch size} < 2550 / M_X$$. Since we can dedicate 3 axes here with a full 3D pod, this would give us a lower bound of 850, which we’re well below. So the answer is no, even with 3 axes. We would be solidly communication-bound.
答案: 既然允许做序列 / 上下文并行了,我们能扩展的规模就大得多。先算每设备 batch size。如果做 8960 路 FSDP,最终每颗 TPU 的 batch size 是 4 * 1024 * 1024 / 8960 = 468 tokens。从上一节我们知道,当 $$\text{每设备 batch size} < 2550 / M_X$$ 时,FSDP 会让我们变成 ICI 受限。既然在一个完整的 3D pod 上我们可以用满 3 条轴,这给出的下界是 850,而我们远低于它。所以答案是:不行,即使有 3 条轴也不行。我们会稳稳地受通信限制。
Question: Now let’s look at mixed tensor parallelism and FSDP. Does there exist some combination that lets us remain compute-bound? What amount of FSDP and tensor parallelism should we do if so?
问题: 现在我们看混合张量并行与 FSDP。是否存在某种组合能让我们保持计算受限?如果可以,FSDP 和张量并行各应该用多少?
Click here for the answer, once you've thought about it!
想好之后,点击这里看答案!
Answer: First let’s check to see if this will even fit. We know that we’ll be comms-bound if our per-chip batch size is less than $2550^2 / 2F = 113$. As we saw above, we’re slightly above this. So that’s great! Now to pick the optimal amount of FSDP, we can use the formula
答案: 先看它到底能不能装下。我们知道,当每芯片 batch size 小于 $2550^2 / 2F = 113$ 时会通信受限。如上面所见,我们略高于这个值。太好了!现在要选最优的 FSDP 用量,可以用公式
$$X_{opt} = \sqrt{\frac{2BN}{F}} = \sqrt{\frac{2 \cdot 4.19e6 \cdot 8960}{28672}} = 1618$$
Rounding to a reasonable multiple of 2, that gives us roughly 2048-way FSDP and 4-way tensor parallelism. That should work well!
取整到 2 的合理倍数,得到大约 2048 路 FSDP 和 4 路张量并行。这应该能跑得不错!
Takeaways: We can train LLaMA-3 with a 4M token batch size on a full TPU v5p pod with a mixture of data parallelism (1024-way), sequence parallelism (2-way), and tensor parallelism (4-way) without being communication-bound. We will be comms-bound if we try to do pure FSDP or FSDP + sequence parallelism. The equations we’ve cooked up in the previous section are very practical.
要点: 我们可以在一个完整的 TPU v5p pod 上,用 4M token 的 batch size 训练 LLaMA-3:混合使用数据并行(1024 路)、序列并行(2 路)和张量并行(4 路),而不会通信受限。如果尝试纯 FSDP 或 FSDP + 序列并行,就会通信受限。上一节推导出的那些公式非常实用。
Worked Problems
实战习题
Question 1 [Scaling LLaMA 70B to more chips]: say we want to train LLaMA 3-70B on 4 pods with the same batch size. What parallelism scheme would we use? Would we be compute or communication bound? Roughly how long would it take to train? Make sure to use the correct roofline bound.
问题 1 [把 LLaMA 70B 扩展到更多芯片]: 假设我们要用相同的 batch size 在 4 个 pod 上训练 LLaMA 3-70B。我们会用哪种并行方案?会是计算受限还是通信受限?训练大约要多久?注意使用正确的 roofline 界限。
Question 2 [LLaMA 405B]:
问题 2 [LLaMA 405B]:
(a) Using the LLaMA 3-405B config (a gated model, so you may need to log in and request access to view it), write a table with all the key hyperparameters as above. How many total parameters does this model have? How many FLOPs per training step? How many FLOPs do we perform if we train for 15T tokens?
(a) 使用 LLaMA 3-405B 的 config(这是一个 gated 模型,你可能需要登录并申请访问权限才能查看),像上面那样列一张包含所有关键超参数的表格。这个模型总共有多少参数?每个训练步多少 FLOPs?如果训练 15T 个 token,总共执行多少次 FLOPs?
(b) Assume we want to train on 8 TPU v5p pods. What parallelism scheme would we use? How long would training take? Would we be compute or comms bound?
(b) 假设我们要在 8 个 TPU v5p pod 上训练。我们会用哪种并行方案?训练要多久?会是计算受限还是通信受限?
讨论
用 GitHub 账号留言;评论保存在公开仓库chengshu-blog-discussions的 Discussions 里。也可通过 RSS 订阅后续文章。