跳转至

计算与参数量

给定序列长度 L、维度 d,推导标准自注意力模块的 FLOPs。

image.png


推导一个标准 Transformer 层的参数量(包括注意力、FFN、归一化)。

以 Transformer Encoder 层为例。

  • 多头自注意力(MHA):

image.png


给定模型层数 L、维度 d、头数 h、词表大小 V,计算 GPT 风格模型的总参数量。

GPT 为 Decoder-only,包含:

  • Token Embedding:V×d

  • L 层 Transformer Decoder,每层含:

image.png

例如,LLaMA-7B:V=32000,d=4096,L=32,dff=11008V=32000,d=4096,L=32,dff=11008 (≈2.7×4096,SwiGLU 三个矩阵),大致计算:嵌入= 32000×4096=131M32000×4096=131M;每层注意力 4×40962=67.1M4×40962=67.1M;FFN 三个矩阵 3×4096×11008=135.2M3×4096×11008=135.2M;总计每层约 202.3M,32层≈6.47B,加上嵌入和解嵌入约 6.74B,接近 7B。


训练一个 7B 模型,FP16 精度,估计模型权重、梯度、优化器状态(Adam)和激活值的显存占用。

假设:7B 模型,参数 FP16(2 bytes),Adam 优化器存储 FP32 的 master weights、动量 m 和 v。

image.png

总显存(粗略):14 + 14 + 84 = 112 GB,再加上激活和其他缓存,远超单卡 80GB,必须使用模型并行和激活检查点。


推理一个 7B 模型,batch size=1,序列长度 2048,估算权重显存和 KV 缓存显存。

  • 权重(FP16):14 GB(如果全部加载到显存)。

image.png

  • 总显存 ≈ 权重 14 GB + KV 1 GB + 其他(输入嵌入、激活、临时缓冲)≈ 16 GB 左右。可在 24GB 显存 GPU 运行。

如果采用 GQA(组数=4),与 MHA 相比 KV 缓存可以减少多少倍?

image.png


在混合精度训练中,forward 和 backward 分别用什么精度?参数 master copy 用什么精度?

标准混合精度训练(AMP):

  • Forward:FP16(或 BF16)计算前向,权重和激活使用半精度,以加速和减少显存。

  • Backward:梯度计算也使用 FP16(或 BF16),但为了避免梯度下溢,损失放大(loss scaling)技术用于 FP16;BF16 因其指数位与 FP32 相同,一般不需要 loss scaling。

  • 参数 master copy:保存 FP32 精度的模型权重,用于累积 FP16 梯度更新,确保精度和数值稳定性。优化器状态(如 Adam 的 m、v)也保持 FP32。

总结:前向/反向计算用 FP16/BF16,主参数和优化器状态用 FP32。


计算 FlashAttention 的 IO 复杂度,并与标准注意力对比。

image.png


一个 Transformer 层中有多少可训练参数矩阵?各是什么形状?

标准 GPT/LLaMA 风格 Decoder 层(Pre-Norm)包含:

  • 自注意力:

image.png

因此,可训练参数矩阵通常为 4 个注意力矩阵 + 3 个 FFN 矩阵 = 7 个矩阵(SwiGLU)或 4+2=6 个矩阵(传统 FFN)。加上 Norm 的向量参数。


如果词表大小为 32k,dmodel=4096,Embedding 层的参数量是多少?LM head 呢?

image.png


分析自回归生成中,生成 T 个 token 的总计算量,与输入长度 L 和生成长度 T 的关系。

在自回归解码中,模型一次生成一个 token,输入序列长度从 L 逐步增长到 L+T-1。对于每一步 t(假设已经生成了 t-1 个 token,当前输入长度为 L + t - 1),模型需要为该序列执行一次前向传播。主要计算量来自自注意力层和 FFN 层。

以标准 Transformer 为例,单个 token 的前向计算量与序列长度有关,因为注意力需要计算当前 token 对所有历史 token 的注意力。对于一个长度为 N 的序列,自注意力的 FLOPs 约为 8Nd2+4N2d(见前文)。但是,利用 KV 缓存后,生成阶段的每一步仅需计算最新 token 的 Q,与完整的 K、V 缓存进行注意力计算。此时,计算量如下:

image.png

因此每生成一个 token,计算量约为 O(d2+Nd)。随着生成的进行,N 从 L 增长到 L+T-1。生成 T 个 token 的总计算量是这些步骤的求和:

image.png

简言之,自回归生成 T 个 token 的总计算量与输入长度 L 成正比(第一项),并且与生成长度的平方 T2T2 成正比(第二项)。实际中由于 KV 缓存,避免了重复计算历史投影,大大加速了推理。


使用梯度检查点后,激活值显存大约可以减少多少倍?重计算的代价是什么?

梯度检查点(Gradient Checkpointing)在训练时仅保存部分层的输入激活,其余中间激活在前向时丢弃,反向传播时再从最近的检查点重新计算前向以恢复所需激活。

  • 减少显存倍数:如果不使用检查点,需要存储所有层的激活,激活显存消耗与层数 L 成正比。使用检查点后,可以只在每 LL 层设置检查点(optimal checkpointing),使显存降低到约 O(L)O(L) 级别。通常对于普通 Transformer,激活显存可以减少 4~10 倍,具体取决于检查点策略。例如,每隔一层做检查点,可减少约一半;更激进的策略可减至 1/10。

  • 重计算代价:反向传播时,需要从最近的检查点重新前向计算被丢弃的层。这相当于增加了一次额外的正向传播。因此,总的计算量大约增加 30%~50%(如果只重算一部分层)。在最坏的情况下(每个层都只存储输入,不存储中间激活),几乎需要重算整个前向,计算量加倍。但在实践中,重计算部分约占总计算时间的 20%~30% 额外开销,却换来了显存的大幅节省,使得更大 batch 或更长序列训练成为可能。


在 8 卡 A100 80GB 上训练一个 65B 模型,可能实现吗?需要哪些并行和 ZeRO 配置?估算。

image.png

  • 张量并行(TP):将单个 Transformer 层的参数切分到多卡,每卡存储部分参数。TP 要求高速通信(NVLink)。8 卡 A100 通常支持 NVLink,可以设置 TP = 8,将每层权重切分到 8 张卡,每卡存储 1/8 参数,即约 16.25 GB。但 TP 通信量较大,且不能无限增加,通常 TP 为 2 或 4。

  • 流水线并行(PP):将不同层放到不同设备上。8 卡可以设置 PP=2, TP=4 或 PP=4, TP=2 等。

  • 数据并行(DP) + ZeRO:使用 ZeRO 优化器将优化器状态、梯度和参数分片到数据并行组。在 8 卡上,如果全局 batch size 较小,可能只有一个数据并行组(即所有卡组成一个模型副本)。此时可以使用 ZeRO-3,将模型参数、梯度和优化器状态都分片到 8 张卡上,每卡存储 1/8。配合 TP/PP 也可以使用 ZeRO。

显存估算(使用 ZeRO-3,无 TP/PP):

  • 参数 FP16:130 GB,分片后每卡 16.25 GB。

  • 梯度 FP16:130 GB,分片后 16.25 GB。

  • 优化器状态 FP32(Adam 需要 12 bytes 每参数):65B×12=780 GB65B×12=780 GB,分片后每卡 97.5 GB。这超过了单卡 80GB,因此单纯 ZeRO-3 无法在 8 卡 A100 上承载 65B(优化器状态过大)。

  • 需要结合 ZeRO-3 + 张量并行 + 流水线并行,或者使用 3D 并行(DP+TP+PP)并可能启用 ZeRO-1/2 辅助。例如,采用 TP=4, PP=2,则模型每卡存储的参数减少 4 倍(TP),2 倍(PP),共 8 倍,与 ZeRO-3 类似。优化器状态也可按 TP/PP 进一步拆分。通常,65B 模型在 8 卡 A100 80GB 上训练是可行的,典型配置如:TP=2, PP=4, DP=1,配合 ZeRO-1(优化器状态分区)或使用 Megatron 的序列并行。或者使用 DeepSpeed ZeRO-3 加上 TP=2 等。具体需要精细估算激活显存和通信开销。

结论:可能,但需要混合并行策略和 ZeRO 优化。例如使用 TP=4, PP=2,然后每张卡上的参数、梯度和优化器状态都会被切分,再结合激活检查点,可以在 8 卡上运行。


给定批大小、序列长度和模型配置,估算每个训练步的前向和反向传播时间。

一个训练步的前向和反向传播时间取决于总 FLOPs 和 GPU 的峰值性能及实际利用率。

image.png

时间估算:

时间 = 总 FLOPs / (GPU 数量 × 单卡峰值 TFLOPs × 硬件利用率)。利用率通常在 30%~60%。例如,A100 FP16 峰值 312 TFLOPS,利用率 0.4,则单卡有效 124.8 TFLOPS。再根据 FLOPs 计算。需要具体模型参数。


如何通过参数量、tokens 数、GPU 数量和利用率估算训练总时间?

image.png


推导 MoE 层中,如果每个 token 激活 top-2 专家,计算量相比 dense 模型增加多少?参数量增加多少?

假设 dense 模型的 FFN 参数量为 2⋅ddff(两个矩阵,如标准 FFN)。MoE 层有 E 个专家,每个专家有相同的 FFN 结构,参数量为 E×2⋅ddff。因此参数量增加了大约 E 倍(实际上与 dense FFN 相比)。

计算量:dense 模型每个 token 计算整个 FFN。MoE 中每个 token 只激活 top-2 专家,所以每个 token 的计算量为 2 个专家的 FFN 运算。因此,每 token 的计算量是 dense 模型的 2 倍(因为 dense 是 1 个 FFN,MoE 计算 2 个专家,若专家与 dense FFN 结构完全相同)。但注意到 dense FFN 的参数量可能比单个专家大?实际上,若为了保持计算量相当,可以设置每个专家的参数量为 dense 的 1/E,但通常 MoE 保持专家大小与 dense FFN 相同,以最大化容量。这种情况下,计算量增加 2 倍,参数量增加 E 倍。总结:MoE 通过稀疏激活,以少量计算开销换取了参数量的大幅增长。


分析 KV 缓存在多轮对话中的增长:10 轮对话,总 token 数 4096,缓存占用是多大?

image.png


为什么在计算注意力复杂度时,通常忽略 softmax 和 mask 操作?它们的开销占比多大?

image.png


如果要把模型部署到手机端,限制内存 2GB,参数量应控制在多少?量化能带来什么帮助?

手机端 2GB 可用内存(含系统和其它应用,假设模型可用 1.5GB)。FP16 精度下每参数 2 字节,1B 参数需要 2GB,因此 FP16 下最多约 1B 参数。但还需要存储推理时的中间激活、KV 缓存等,实际应更小。如果采用 4-bit 量化(如 Q4_K),每个参数约 4 位(0.5 字节),1B 参数约 0.5GB,因此可以支持 2~3B 参数模型。此外,8-bit 量化也能将模型减小一半。量化不仅减少存储,还可能利用整数运算单元加速推理,减少内存带宽压力,使手机端运行更加流畅。结合剪枝、蒸馏等技术,可以在 2GB 内存限制下实现较好的体验。


写一个脚本或伪代码,基于给定模型配置计算总的可训练参数量。

以下 Python 伪代码,计算类似 LLaMA 结构的参数量:

def calc_llama_params(vocab_size, d_model, n_layers, n_heads, ff_dim, shared_emb=True):
    # Embedding
    embed_params = vocab_size * d_model
    # LM head (if not shared)
    lm_head_params = 0 if shared_emb else vocab_size * d_model
    # Per layer
    attn_params = 4 * d_model * d_model  # Q, K, V, O without bias
    # FFN (SwiGLU with 3 matrices)
    ffn_params = 3 * d_model * ff_dim  # gate, up, down
    rmsnorm_params = 2 * d_model  # two RMSNorm scales
    layer_params = attn_params + ffn_params + rmsnorm_params
    total = embed_params + lm_head_params + n_layers * layer_params
    return total

# Example for LLaMA-7B
config = {
    'vocab_size': 32000,
    'd_model': 4096,
    'n_layers': 32,
    'n_heads': 32,
    'ff_dim': 11008,
    'shared_emb': True
}
params = calc_llama_params(**config)
print(f"Total parameters: {params/1e9:.2f}B")

该脚本累加嵌入、层注意力和 FFN 参数,忽略位置编码和可能的偏置,能给出非常接近实际的估计值。


通过以上推导和伪代码,我们全面解答了自回归生成的计算量、梯度检查点、分布式训练可行性、训练时间估算、MoE 开销、KV 缓存增长、softmax 忽略原因、手机端部署限制以及参数量计算脚本。