01

整体架构

GPU 显存消耗分为两大类:

类型包含内容特点
静态显存权重 (weight)、FP8 转置缓存 (fp8_cache)、梯度 (grad)、分片参数 (sharded_param)、优化器状态 (optimizer_state)训练期间固定占用,与 batch size / seq_len 无关
动态显存 (Activation)前向传播中 save_for_backward 保存的张量随 micro_batch_size 和 seq_len 线性增长

计算器的主函数 calculate() 的执行流程:

graph TD A["读取输入参数
(nodes, PP, TP, CP, EP, hidden_size...)"] --> B["计算派生值
(DP, EDP, SEDP, VPP...)"] B --> C["parsePPLayout()
解析 PP 布局字符串"] C --> D["构建 chunk_meta
每个 PP rank 持有的层类型统计"] D --> E["遍历各模块
调用 getModuleMemory()"] E --> F["计算每层 Activation
逐子模块求和"] F --> G["汇总: 静态 + 动态 + NCCL buffer"]

源码参考:megatron/training/theoretical_memory_usage.py 提供了官方的理论估算函数 compute_weight_and_optimizer_memory()compute_activation_memory()。计算器在此基础上做了更精细的逐模块拆分。

02

并行维度与分片大小

2.1 并行度计算

total_gpus = nodes * 8                    # 每节点 8 卡

DP   = total_gpus / PP / TP / CP          # 数据并行度
EDP  = total_gpus / PP / EP               # Expert 数据并行度
SEDP = total_gpus / PP                    # Shared Expert 数据并行度 (= TP * CP * DP)
VDP  = total_gpus                         # ViT 数据并行度 (全局)

2.2 shard_size 的含义

使用 Distributed Optimizer (ZeRO-1) 时,优化器状态和 master weights 按 DP rank 分片。不同模块因并行方式不同,其 shard_size 也不同:

模块shard_size原因
Embedding, AttentionQKV/Out, SharedExpert, OutputLayerDP已被 TP 切分,TP group 内每卡持有不同参数,只在 DP 维度分片
MoE Gate, DenseFC1/FC2SEDP (= TP×CP×DP)不做 TP 切分,所有 GPU 持有相同副本,分片范围更大
MoE Experts (FC1/FC2)EDPExpert 只在 EP 组内复制,分片在 Expert DP 维度

源码对应megatron/core/optimizer/distrib_optimizer.py 中,每个参数组根据其所属的 process group 确定分片大小。参数按 ParamAndGradBuffer 中的 bucket 划分,每个 bucket 的 DP group 决定了 shard_size。

graph LR subgraph "TP 切分的模块 (Attention, Embedding...)" A1["GPU 0: 参数片 A"] --- A2["GPU 1: 参数片 B"] A3["GPU 2: 参数片 A"] --- A4["GPU 3: 参数片 B"] end subgraph "DP 分片优化器状态" A1 -.- A3 A2 -.- A4 end

TP=2, DP=2: TP 内持有不同参数片,DP 内持有相同参数片 → 优化器状态只需在 DP=2 维度分片。

03

getModuleMemory() 静态显存公式

核心函数,为每个模块计算五项显存(单位 GiB):

def getModuleMemory(num_elements, element_size, shard_size, has_fp8_cache=False):
    # FP8 权重可选择存为 2 bytes (BF16 副本)
    weight_bytes = 2 if (weight_2b_for_1b and element_size < 2) else element_size

    weight        = num_elements * weight_bytes / (1024**3)            # 权重本体
    fp8_cache     = weight if has_fp8_cache else 0                     # FP8 转置缓存
    grad          = num_elements * 4 / (1024**3)                       # 梯度 (固定 FP32)
    sharded_param = num_elements * 4 / shard_size / (1024**3)          # Master weights 分片
    optimizer_state = num_elements * 4 / shard_size / (1024**3)        # 优化器动量分片

3.1 各项详解

weight(权重本体)

  • BF16 训练:element_size = 2,每个参数 2 bytes
  • FP8 训练:element_size = 1,但若 weight_2b_for_1b = true,实际以 BF16 (2 bytes) 存储
  • 源码:megatron/core/tensor_parallel/layers.py — weight 的 dtype 由 params_dtype 决定

fp8_cache(FP8 转置缓存)

  • FP8 训练中,反向传播需要权重的转置版本。为避免反复计算,缓存转置后的权重
  • 仅线性层有此缓存(Embedding / Gate / OutputLayer 无)
  • 大小等于 weight 本身
  • 源码:megatron/core/fp8_utils.py — lazy transposition cache

grad(梯度)

  • 固定 FP32 (4 bytes),不论前向精度如何
  • 梯度在反向传播时以 torch.float32 累积,确保数值精度

sharded_param(Master Weights 分片)

  • Distributed Optimizer 中,FP32 master weights 被均匀分片到各 DP rank
  • 每个 rank 只存 num_elements / shard_size 个参数的 FP32 副本
  • 源码:distrib_optimizer.py_build_model_gbuf_range_map() 方法完成分片映射

optimizer_state(优化器状态)

  • Adam 优化器包含:momentum (FP32) + variance (FP32) = 8 bytes/element
  • 分片后每 rank:num_elements * 8 / shard_size bytes
  • 计算器简化为一个字段 (4 bytes),对应 momentum 的一阶矩

3.2 源码对照:theoretical_memory_usage.py

# 源码中的等价公式 (line 179-183)
num_bytes_per_parameter = (
    18 if not args.use_distributed_optimizer
    else 6 + (12 / args.data_parallel_size)
)
# 18 = 2(weight) + 4(grad) + 4(master) + 4(momentum) + 4(variance)
# 分布式优化器: 6 = 2(weight) + 4(grad), 12/DP = (4+4+4)/DP
04

各模块参数量

以下公式为每个 PP rank 上单层或单模块的参数元素数量(num_elements),传入 getModuleMemory() 后乘以该 PP rank 上的层数。

4.1 Embedding

num_embedding_params = vocab_size * hidden_size / TP
  • 词表沿 vocab 维度按 TP 切分(VocabParallelEmbedding
  • shard_size = DPhas_fp8_cache = false
  • 源码:megatron/core/models/gpt/gpt_embedding.py

4.2 Attention QKV (ColumnParallelLinear)

num_LinearQKV_params = kv_channels * (num_attn_heads * (1 + output_gate) + 2 * num_query_groups) * hidden_size / TP
  • kv_channels = head_dim(每个注意力头的维度)
  • num_attn_heads = Q 头数
  • num_query_groups = KV 头数(GQA 时 < num_attn_heads)
  • output_gate:是否使用 attention output gate(额外一组 Q 投影)
  • QKV 融合为一个 ColumnParallelLinear,output 维度 = Q + K + V 的 head dims 之和
  • 源码:attention.pyself.linear_qkv
  • 源码:layers.pyColumnParallelLinear 沿 output 维度按 TP 切分

4.3 Attention Output Projection (RowParallelLinear)

num_LinearOut_params = kv_channels * num_attn_heads * hidden_size / TP
  • 权重形状 [query_projection_size / TP, hidden_size]
  • 源码:attention.pyself.linear_projRowParallelLinear,input 维度按 TP 切分

4.4 MoE Gate (Router)

num_MoE_Gate_params = hidden_size * num_experts
  • 不做 TP 切分,每个 GPU 都有完整副本
  • shard_size = SEDP(所有 TP/CP rank 持有相同参数,可在更大范围分片优化器状态)
  • 源码:megatron/core/transformer/moe/router.pyTopKRouter 中的 self.weight

4.5 MoE Shared Expert FC1

num_MoE_LinearFC1_shared_expert_params = shared_expert_intermediate_size * hidden_size * 2 * shared_expert / TP
  • * 2:SwiGLU 门控导致 FC1 输出维度翻倍(gate 分支 + up 分支)
  • shared_expert:0 或 1 (是否启用共享专家)
  • 按 TP 切分,shard_size = DP
  • 源码:moe_layer.pyself.shared_experts 使用标准 MLP 模块

4.6 MoE Routed Expert FC1

num_local_experts = num_experts / EP
num_MoE_LinearFC1_experts_params = num_local_experts * hidden_size * ffn_hidden_size * 2
  • num_local_experts:每个 EP rank 上的本地专家数
  • 不做 TP 切分GroupedMLP 内部以 Grouped GEMM 整体执行)
  • shard_size = EDP
  • 源码:experts.py:128-133fc1_output_size = moe_ffn_hidden_size * num_local_experts (* 2 if gated)

4.7 MoE FC2 (Shared + Routed)

# Shared Expert FC2
num_MoE_LinearFC2_shared_expert_params = shared_expert_intermediate_size * hidden_size * shared_expert / TP

# Routed Expert FC2
num_MoE_LinearFC2_experts_params = num_local_experts * hidden_size * ffn_hidden_size
  • FC2 无门控因子(SwiGLU 的 ×2 只影响 FC1)

4.8 Dense Layer FC1 / FC2

dense_ffn_hidden_size = ffn_hidden_size * topk

num_dense_LinearFC1_params = dense_ffn_hidden_size * hidden_size * 2   # * 2 for SwiGLU
num_dense_LinearFC2_params = dense_ffn_hidden_size * hidden_size
  • Dense 层不做 TP 切分,shard_size = SEDP
  • 源码:megatron/core/transformer/mlp.py

4.9 Output Layer

num_OutputLayer_params = vocab_size * hidden_size / TP
  • 与 Embedding 可共享权重(untie_embeddings_and_output_weights
  • 源码:megatron/core/models/gpt/gpt_model.pyself.output_layer

4.10 模块汇总表

模块num_elementselement_sizeshard_sizefp8_cache
Embeddingvocab × hidden / TPemb_bytesDPNo
AttentionQKVkv_ch × (heads×(1+gate) + 2×kv_groups) × hidden / TPlinear_bytesDPYes
AttentionOutkv_ch × heads × hidden / TPlinear_bytesDPYes
Gatehidden × num_expertslogits_bytesSEDPNo
MoEFC1/sharedshared_ffn × hidden × 2 / TPlinear_bytesDPYes
MoEFC1/expertslocal_experts × hidden × ffn × 2linear_bytesEDPYes
MoEFC2/sharedshared_ffn × hidden / TPlinear_bytesDPYes
MoEFC2/expertslocal_experts × hidden × ffnlinear_bytesEDPYes
DenseFC1ffn×topk × hidden × 2linear_bytesSEDPYes
DenseFC2ffn×topk × hiddenlinear_bytesSEDPYes
OutputLayervocab × hidden / TPoutput_bytesDPNo
05

PP Layout 与 VPP

5.1 PP Layout 语法

parsePPLayout() 解析 Python 风格的列表表达式,例如:

['Ett'] + ['ttt'] * 18 + ['mL']

解析结果是一个字符串数组,每个元素代表一个 PP stage 的 model chunk。每个字符表示一种层类型:

字符含义
EEmbedding 层
tTransformer 层(按模型顺序分配 MoE 或 Dense)
mMTP 层 (Multi-Token Prediction)
LOutput Layer(logits / loss 计算)

5.2 VPP 计算

num_stage = len(pp_layout)    # PP layout 中的总 stage 数
VPP = num_stage / PP          # 每个 PP rank 持有的 model chunk 数

例如 PP=4,20 个 stage → VPP=5,每个 PP rank 执行 5 个 model chunks。

5.3 chunk_meta 构建

计算器遍历 PP layout,为每个 PP rank 统计其负责的各类层数量:

# 对每个 PP rank r (0 到 PP-1):
for stage in range(r, num_stage, PP):
    chunk = pp_layout[stage]
    for char in chunk:
        if char == 'E': count_Embedding += 1
        elif char == 't':
            transformer_idx += 1
            if transformer_idx <= num_dense_layers:
                count_Dense += 1
            else:
                count_MoE += 1
        elif char == 'm': count_MTP += 1
        elif char == 'L': count_OutputLayer += 1

t 按模型中出现的顺序依次分配类型:前 num_dense_layers 个标记为 Dense,之后标记为 MoE。

5.4 PP rank 的参数量

每个 PP rank 的总参数量 = 其持有的各类层数量 × 每层的参数量。

# 计算器对每个 PP rank 调用 getModuleMemory():
modules = [
    ['Embedding',     chunk_meta.Embedding * num_embedding_params,     emb_bytes,  DP,   False],
    ['AttentionQKV',  num_chunk_layers * num_LinearQKV_params,         lin_bytes,  DP,   True],
    ['AttentionOut',  num_chunk_layers * num_LinearOut_params,         lin_bytes,  DP,   True],
    ['Gate',          chunk_meta.MoELayer * num_MoE_Gate_params,       log_bytes,  SEDP, False],
    ['MoEFC1/shared', chunk_meta.MoELayer * num_MoE_FC1_shared,       lin_bytes,  DP,   True],
    ['MoEFC1/experts',chunk_meta.MoELayer * num_MoE_FC1_experts,      lin_bytes,  EDP,  True],
    # ... (MoEFC2, DenseFC1, DenseFC2, OutputLayer 类似)
]
06

Activation 动态显存

6.1 基本单位

所有 activation 公式的基本形式(单位 MB):

activation_MB = intDiv(intDiv(mbs * seqlen * dim * dtype_bytes, TP), CP) / (1024**2)
  • intDiv(x, y) = floor(x / y):整数除法,与实际 tensor shape 计算一致
  • mbs:micro batch size
  • seqlen:序列长度
  • dim:该张量的特征维度
  • / TP:张量并行切分(特征维度)
  • / CP:上下文并行切分(序列维度)

6.2 Transformer 层前向流程与 Activation

每个 Transformer 层的前向传播流程(源码 transformer_layer.py):

graph TD A["Input hidden_states"] --> B["Pre-Attn LayerNorm"] B --> C["Linear QKV"] C --> D["Core Attention
(Q,K,V → O)"] D --> E["Attention Output Proj"] E --> F["Bias-Dropout-Add + Residual"] F --> G["Pre-MLP LayerNorm"] G --> H{"MoE or Dense?"} H -- MoE --> I["Gate/Router"] I --> J["Shared Expert"] I --> K["Routed Experts"] J --> L["Combine"] K --> L H -- Dense --> M["Dense MLP"] L --> N["Bias-Dropout-Add + Residual"] M --> N

6.3 逐子模块 Activation 公式

(a) Pre-Attention LayerNorm

pre_attn_layernorm = mbs * seqlen * hidden_size * 2 / TP / CP    # BF16 输入

保存 LayerNorm 的输入用于反向传播。可选重计算:transformer_layer.py:531-536CheckpointWithoutOutput()

(b) Linear QKV (ColumnParallelLinear)

linear_qkv = mbs * seqlen * hidden_size * linear_bytes / TP / CP
  • linear_bytes = 2 (BF16) 或 1 (FP8)
  • 保存输入 hidden_states 用于计算权重梯度
  • 源码:layers.py:481ctx.save_for_backward(input, weight)

(c) Core Attention (Q, K, V, O)

attention_q = mbs * seqlen * kv_channels * num_attn_heads * 2 / TP / CP
attention_k = mbs * seqlen * kv_channels * num_query_groups * 2 / TP / CP
attention_v = attention_k
attention_o = attention_q    # attention output,与 Q 同形状
  • FlashAttention 中 save_for_backward(q, k, v, ..., *outs)
  • 源码:dot_product_attention_context_parallel.pyctx.save_for_backward(q, k, v, attention_mask, *outs, *probs)
  • Q 和 O 按 num_attn_heads / TP 切分;K, V 按 num_query_groups / TP 切分
  • 均为 BF16 (2 bytes)

(d) Attention Output Projection

RowParallelLinear 保存 core_attn_out 作为输入。FP8 时额外保存 BF16 原始输入:

# attention.py:266-268
set_save_original_input(self.linear_proj)  # FP8 时保存量化前的 BF16 输入

(e) Pre-MLP LayerNorm

pre_mlp_layernorm = mbs * seqlen * hidden_size * 2 / TP / CP

(f) MoE Gate Activation

gate_act = mbs * seqlen * hidden_size * 2 / TP / CP
  • Router 的 RouterGatingLinearFunction.save_for_backward(inp, weight, bias)
  • 保存 BF16 hidden_states 用于 gate 权重的梯度计算
  • 源码:moe_utils.py

(g) Shared Expert Activations(每 MoE 层)

# SwiGLU clamp 操作保存的 mask (1 byte per element)
shared_clamp     = mbs * seqlen * shared_intermediate * 2 * 1 / TP / CP * shared_expert

# SwiGLU 两个分支的激活值 (gate + up), 各 2 bytes
shared_activation = mbs * seqlen * shared_intermediate * 2 * 2 / TP / CP * shared_expert

# FC2 输入(经过激活后的中间结果)
shared_fc2       = mbs * seqlen * shared_intermediate * linear_bytes / TP / CP * shared_expert

(h) MoE Routed Expert Activations(每 MoE 层)

# FC1 输入(dispatched hidden states)
moe_fc1        = mbs * seqlen * topk * droprate * hidden_size * linear_bytes / TP / CP

# SwiGLU clamp mask
moe_clamp      = mbs * seqlen * topk * droprate * ffn_hidden_size * 2 * 1 / TP / CP

# SwiGLU 激活值
moe_activation = mbs * seqlen * topk * droprate * ffn_hidden_size * 2 * 2 / TP / CP

# FC2 输入
moe_fc2        = mbs * seqlen * topk * droprate * ffn_hidden_size * linear_bytes / TP / CP
  • topk * droprate:每个 token 被路由到 topk 个专家,droprate 是 token drop 后的保留率
  • 源码:experts.py:261-274gg.ops.gmm 的 grouped GEMM 操作

(i) Dense Layer Activations(每 Dense 层)

dense_ffn_hidden_size = ffn_hidden_size * topk

dense_fc1        = mbs * seqlen * hidden_size * linear_bytes / TP / CP
dense_clamp      = mbs * seqlen * dense_ffn_hidden_size * 2 * 1 / TP / CP
dense_activation = mbs * seqlen * dense_ffn_hidden_size * 2 * 2 / TP / CP
dense_fc2        = mbs * seqlen * dense_ffn_hidden_size * linear_bytes / TP / CP
  • 结构与 MoE Expert 类似,但使用 dense_ffn_hidden_size,无 routing 开销

(j) Output Layer Activation

output_layer_act = mbs * seqlen * vocab_size * 2 * 3 / TP / CP
  • * 3:Cross-entropy loss 保存三个张量
  • 源码:cross_entropy.py:187ctx.save_for_backward(exp_logits, target_mask, masked_target_1d)
  • exp_logits(softmax 结果)是最大的,形状 [mbs*seqlen, vocab/TP]

6.4 每层 Activation 汇总

# MoE 层 activation
moe_layer_act = (pre_attn_layernorm + linear_qkv
    + attention_q + attention_k + attention_v + attention_o
    + pre_mlp_layernorm + gate_act
    + shared_clamp + shared_activation + shared_fc2
    + moe_fc1 + moe_clamp + moe_activation + moe_fc2)

# Dense 层 activation
dense_layer_act = (pre_attn_layernorm + linear_qkv
    + attention_q + attention_k + attention_v + attention_o
    + pre_mlp_layernorm
    + dense_fc1 + dense_clamp + dense_activation + dense_fc2)

# PP rank 总 activation
rank_activation = (sum_moe_layers + sum_dense_layers) * VPP + output_layer_act
07

MoE Activation Recomputation

7.1 触发条件

# experts.py:102-105
self.activation_recompute = (
    self.config.recompute_granularity == 'selective'
    and "moe_act" in self.config.recompute_modules
)

recompute_granularity = 'selective'recompute_modules 包含 "moe_act" 时启用。

7.2 重计算机制

# experts.py:264-269
if self.activation_recompute:
    # CheckpointWithoutOutput: 只保存输入,丢弃激活函数的输出
    intermediate_parallel = self.activation_checkpoint.checkpoint(
        self.activation_func_with_probs, fc1_output, permuted_probs.unsqueeze(-1)
    )
    fc2_output = gg.ops.gmm(intermediate_parallel, w2, tokens_per_expert, trans_b=False)
    # 丢弃 fc2_output 的保存,反向传播时重新计算
    self.activation_checkpoint.discard_output_and_register_recompute(fc2_output)
graph LR subgraph "正常前向(无重计算)" A1["FC1 output"] --> B1["SwiGLU 激活"] B1 --> C1["FC2"] B1 -. "save clamp" .-> S1["saved"] B1 -. "save activation" .-> S2["saved"] C1 -. "save fc2 input" .-> S3["saved"] end subgraph "moe_act 重计算" A2["FC1 output"] --> B2["SwiGLU 激活"] B2 --> C2["FC2"] B2 -. "discard" .-> X1["freed"] C2 -. "discard" .-> X2["freed"] end

7.3 显存节省量

启用 moe_act 重计算后,从每个 MoE 层的 activation 中去除三项:

savings_per_moe_layer = moe_clamp + moe_activation + moe_fc2

以 hidden=7168, ffn=2048, topk=8, seqlen=4096, mbs=1, TP=8, CP=1 为例:

moe_clamp      = 1 * 4096 * 8 * 2048 * 2 * 1 / 8 = 16 MB
moe_activation = 1 * 4096 * 8 * 2048 * 2 * 2 / 8 = 32 MB
moe_fc2        = 1 * 4096 * 8 * 2048 * 1 / 8     =  8 MB
# 每 MoE 层节省约 56 MB (FP8) 或更多 (BF16)

7.4 代价

反向传播时需要重新计算 SwiGLU 激活函数和 FC2 的前向,增加约 30% 的 MoE 计算量。这是经典的 计算换内存 tradeoff。

08

VPP 对 Activation 的影响

8.1 为什么 VPP 增加 Activation 占用?

VPP (Virtual Pipeline Parallelism) 将每个 PP rank 的层分成多个 model chunks。在 Interleaved Schedule 中,多个 microbatch 同时在不同 chunk 上执行,每个 chunk 都需要保留其 in-flight microbatch 的 activation。

8.2 计算器的乘数

# 计算器使用 VPP 作为直接乘数(保守估计)
total_activation = per_chunk_activation * VPP

8.3 源码的乘数

# theoretical_memory_usage.py:224-227
interleaved_schedule_memory_penalty = 1 + (PP - 1) / (PP * VPP)

8.4 对比

PPVPP计算器乘数源码乘数
4111.75
4221.375
4551.15
8441.21875

计算器的 VPP 乘数是保守上界:假设每个 chunk 都有独立的 microbatch activation 同时存在。源码的公式更精确,反映了实际调度中 in-flight microbatch 的数量。

实际 peak activation 取决于调度细节(warmup / steady / cooldown 阶段的 microbatch 重叠情况),两者各有适用场景。

09

ViT 显存

计算器单独处理 Vision Transformer (ViT) 的参数和显存。ViT 不做 TP/PP 切分,所有参数用 VDP(全局 GPU 数)分片。

9.1 ViT 参数公式

# Patch Embedding (Conv2d)
vit_conv_params = patch_dim**2 * 3 * vit_hidden_size

# Self-Attention(所有 ViT 层)
vit_QKV_params = vit_num_layers * vit_kv_channels * (vit_num_attn_heads + 2 * vit_num_query_groups) * vit_hidden_size
vit_Out_params = vit_num_layers * vit_kv_channels * vit_num_attn_heads * vit_hidden_size

# MLP(所有 ViT 层)
vit_FC1_params = vit_num_layers * vit_ffn_hidden_size * vit_hidden_size
vit_FC2_params = vit_num_layers * vit_ffn_hidden_size * vit_hidden_size

9.2 Projector & Merge 层

# 视觉-语言投影层
projector_FC1_params = vit_hidden_size * projector_ffn_hidden_size
projector_FC2_params = projector_ffn_hidden_size * hidden_size

# 空间合并层
merge_FC1_params = projector_ffn_hidden_size * hidden_size * spatial_merge_size**2
merge_FC2_params = projector_ffn_hidden_size * hidden_size

9.3 ViT 显存特点

  • ViT 以 BF16 训练(element_size = 2),无 FP8
  • shard_size = VDP(全局所有 GPU),因 ViT 不参与 PP/TP/EP
  • ViT 的 activation 独立计算,不计入 LLM 的 VPP 乘数
10

总显存汇总公式

10.1 端到端计算流程

graph TD A["1. 解析 PP Layout"] --> B["2. 确定每个 PP rank 的层组成"] B --> C["3. 计算每模块 num_elements"] C --> D["4. getModuleMemory × 层数"] D --> E["5. 累加所有模块的静态显存"] B --> F["6. 计算每层 Activation"] F --> G["7. 各层 Activation × VPP + OutputLayer"] E --> H["8. 总显存 = 静态 + 动态 + NCCL buffer"] G --> H

10.2 汇总公式

# 静态显存 (每个 PP rank)
static_memory = sum(
    getModuleMemory(chunk_layers * module_params, elem_size, shard_size, fp8_cache)
    for module in [Embedding, QKV, Out, Gate, SharedFC1, ExpertFC1,
                   SharedFC2, ExpertFC2, DenseFC1, DenseFC2, OutputLayer]
)

# 动态显存 (每个 PP rank)
per_moe_layer = (pre_attn_ln + linear_qkv + Q + K + V + O
    + pre_mlp_ln + gate
    + shared_clamp + shared_act + shared_fc2
    + moe_fc1 + moe_clamp + moe_act + moe_fc2)

per_dense_layer = (pre_attn_ln + linear_qkv + Q + K + V + O
    + pre_mlp_ln
    + dense_fc1 + dense_clamp + dense_act + dense_fc2)

dynamic_memory = (num_moe_layers * per_moe_layer
    + num_dense_layers * per_dense_layer) * VPP
    + output_layer_act

# 总显存
total_gpu_memory = static_memory + dynamic_memory + nccl_buffer

10.3 源码中的等价公式

theoretical_memory_usage.py 提供简化版公式作为参照:

# 权重 + 优化器
weight_and_optimizer = num_params_on_shard * bytes_per_param
# 其中 bytes_per_param = 18 (无分布式优化器) 或 6 + 12/DP (有分布式优化器)

# Activation (每层)
activation_per_layer = seq_len * mbs * hidden * (18 + 4 * (ffn / hidden))
# 18 包括: layernorm(2) + QKV(2) + Q(2) + K(2) + V(2) + O(2) + proj(2) + mlp_ln(2) + fc1(2)
# 4*(ffn/hidden) 包括: clamp(1) + activation(2) + fc2(1) 按 ffn 缩放

# VPP penalty
activation *= 1 + (PP - 1) / (PP * VPP)

10.4 计算器 vs 源码的差异

维度计算器源码 (theoretical_memory_usage.py)
精度intDiv 整数除法,匹配实际 tensor shape浮点除法
优化器逐模块 sharded_param + optimizer_state186 + 12/DP 整体估算
MoE/Dense 区分分别计算,支持混合 PP layout不区分 MoE/Dense 层
VPP 乘数VPP(保守上界)1 + (PP-1)/(PP*VPP)(更精确)
FP8 支持区分 linear_bytes, fp8_cache不支持 FP8
MoE 重计算可选去除 clamp + activation + fc2不支持
ViT独立计算 ViT + Projector不支持
Cross-Entropyvocab * 2 * 3 近似hidden * 4 * (1 + vocab/hidden)