| # 面试:Transformer 架构深度 |
|
|
| > 本仓库 `src/core/` 的纯 Transformer 组件,对应代码:`norm.py`、`rope.py`、`attention.py`、`mlp.py`、`block.py` |
|
|
| ## 0. 整体架构图 |
|
|
| ``` |
| 输入 token id |
| │ |
| ▼ |
| ┌─────────────┐ |
| │ TokenEmbed │ (vocab_size × d_model) |
| └──────┬──────┘ |
| │ + 位置编码(RoPE 作用于 Q/K) |
| ▼ |
| ┌──────────────────────────────┐ |
| │ Transformer Block × N │ |
| │ ┌────────────────────────┐ │ |
| │ │ RMSNorm │ │ |
| │ │ MultiHeadSelfAttention │ │ ← Q/K 先过 QK-Norm → RoPE |
| │ │ SDPA (causal mask) │ │ |
| │ └───────────┬────────────┘ │ |
| │ 残差相加 │ │ |
| │ ┌────────────▼───────────┐ │ |
| │ │ RMSNorm │ │ |
| │ │ SwiGLU FFN / MoE │ │ ← config.use_moe 切换 |
| │ └───────────┬────────────┘ │ |
| │ 残差相加 │ │ |
| └──────────────┼───────────────┘ |
| │ |
| ▼ |
| RMSNorm → lm_head → logits (vocab_size) |
| ``` |
|
|
| > 本仓库 `LM`(`src/models/lm/model.py:11`):TokenEmbed → N×Block → RMSNorm → lm_head |
| |
| --- |
| |
| ## Q1. RMSNorm 和 LayerNorm 区别?为什么选 RMSNorm? |
| |
| ### 公式对比 |
| |
| | | LayerNorm | RMSNorm | |
| |---|---|---| |
| | 公式 | `(x - mean) / std * γ + β` | `x / sqrt(mean(x²) + eps) * γ` | |
| | 减均值 | ✅ 是 | ❌ 否 | |
| | 可学习参数 | γ, β | γ | |
| |
| ### 本仓库实现(`src/core/norm.py:5-15`) |
| |
| ```python |
| class RMSNorm(nn.Module): |
| def forward(self, x): |
| x = x.float() # 先转 float32 保精度 |
| x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) |
| return (self.weight * x).type_as(self.weight) # 再转回原 dtype |
| ``` |
| |
| ### 为什么 RMSNorm 更好? |
| |
| 1. **计算效率**:省去均值计算,减少一次归一化操作 |
| 2. **训练稳定性**:LLaMA/GPT-NeoX 系列标配,实践证明足够 |
| 3. **精度技巧**:先转 float32 计算,再转回原 dtype——混合精度下的标准做法 |
| |
| > 面试点:为什么先转 float32?→ 避免 float16/bf16 下的数值溢出 |
| |
| --- |
| |
| ## Q2. RoPE 是什么?为什么好?(`src/core/rope.py`) |
| |
| ### 公式推导 |
| |
| 把 d_model 维向量按相邻两维配对,对第 m 个位置、第 i 对 (2i, 2i+1) 施加旋转角 θ_i·m: |
| |
| ``` |
| θ_i = base^(-2i/d) # base 即 rope_theta,默认 1e6 |
|
|
| [cos(θ_i·m) -sin(θ_i·m)] [x_{2i} ] |
| [sin(θ_i·m) cos(θ_i·m)] · [x_{2i+1}] |
| ``` |
| |
| ### 为什么编码的是相对位置 |
| |
| 旋转矩阵是正交阵,关键是**可加性**:旋转角相加 = 位置差。 |
| |
| ``` |
| (R_m q)ᵀ (R_n k) = qᵀ R_{n-m} k |
| ``` |
| |
| 即 Q@K 内积只依赖 **(n-m)** 这个相对位置 → 平移不变、长度外推好。 |
| |
| ### 本仓库实现(`src/core/rope.py:5-37`) |
| |
| ```python |
| def precompute_freqs_cis(dim, max_position_embeddings=32768, base=1e6, rope_scaling=None): |
| # 基础频率计算 |
| inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim)) |
| |
| # YaRN 扩展(可选) |
| if rope_scaling is not None: |
| # beta_fast=32, beta_slow=1, factor=16 |
| # 中频段频率缩放 + 注意力因子修正 |
| pass |
| |
| t = torch.arange(max_position_embeddings) |
| freqs = torch.outer(t, inv_freq) |
| |
| # 复制一份用于 Q/K |
| freqs_cos = torch.cat([freqs, freqs], dim=-1) |
| freqs_sin = torch.cat([freqs, freqs], dim=-1) |
| return freqs_cos, freqs_sin |
| |
| def apply_rotary_pos_emb(q, k, freqs_cos, freqs_sin): |
| # 半维翻转 + 旋转 |
| q_rot = rotate_half(q) |
| k_rot = rotate_half(k) |
| q = q * freqs_cos + q_rot * freqs_sin |
| k = k * freqs_cos + k_rot * freqs_sin |
| return q.type_as(q_orig), k.type_as(k_orig) |
| ``` |
| |
| ### 对比:RoPE vs 绝对位置 Embedding |
|
|
| | | RoPE | 绝对位置 Embedding | |
| |---|---|---| |
| | 内积是否含相对位置 | ✅ 自动 | ❌ 需学 | |
| | 长度外推 | 好(可调整 base) | 差(超出训练长度即 OOV) | |
| | 额外参数 | 无 | 有 (context_length×d) | |
| | 代表模型 | LLaMA, GPT-NeoX | GPT-2, BERT | |
| |
| ### YaRN:长上下文扩展 |
| |
| 当 `rope_scaling` 不为 None 时启用,支持: |
| - `factor=16`:上下文长度扩展倍数 |
| - `original_max_position_embeddings=2048`:原始训练长度 |
| - 中频段频率做 `1/factor` 缩放 + 注意力因子修正(NTK-aware) |
|
|
| > 面试点:YaRN 为什么比直接外推 RoPE 更好?→ 直接外推会导致高频位置编码失效,YaRN 通过分频段处理解决这个问题 |
|
|
| --- |
|
|
| ## Q3. GQA 解决了什么?(`src/core/attention.py:10-55`) |
|
|
| ### 三种注意力模式 |
|
|
| | | MHA | GQA | MQA | |
| |---|---|---|---| |
| | Q 头数 | n_heads | n_heads | n_heads | |
| | KV 头数 | n_heads | n_kv_heads | 1 | |
| | 显存 | 高 | 中 | 低 | |
| | 质量 | 高 | 中 | 低 | |
| | 代表 | GPT-3 | LLaMA-2 | PaLM | |
|
|
| ### 本仓库实现(`src/core/attention.py:13-16`) |
|
|
| ```python |
| class Attention(nn.Module): |
| def __init__(self, config): |
| self.n_local_heads = config.num_attention_heads # 8 |
| self.n_local_kv_heads = config.num_key_value_heads # 4 |
| self.n_rep = self.n_local_heads // self.n_local_kv_heads # 2 倍复制 |
| ``` |
|
|
| ### repeat_kv 实现(`src/core/rope.py:33-37`) |
| |
| ```python |
| def repeat_kv(x, n_rep): |
| # (bs, slen, num_kv_heads, head_dim) -> (bs, slen, num_heads, head_dim) |
| if n_rep == 1: |
| return x |
| return x[:, :, :, None, :].expand(bs, slen, num_kv_heads, n_rep, head_dim).reshape(bs, slen, num_heads, head_dim) |
| ``` |
| |
| > 面试点:GQA 的 KV cache 节省多少?→ 假设 8 Q 头 / 4 KV 头,KV cache 减半,推理显存显著降低 |
|
|
| --- |
|
|
| ## Q4. QK-Norm 为什么有用?(`src/core/attention.py:23-24, 36`) |
|
|
| ### 问题:Attention Logit 爆炸 |
|
|
| 在深层 Transformer 中,Q/K 的 norm 可能会越来越大,导致: |
| - softmax 进入饱和区 |
| - 梯度消失 |
| - 训练不稳定 |
|
|
| ### 解决方案:在 RoPE 前做 RMSNorm |
|
|
| ```python |
| # src/core/attention.py:36 |
| q = self.q_norm(self.q_proj(x)) # 先做 QK-Norm |
| k = self.k_norm(self.k_proj(x)) |
| q, k = apply_rotary_pos_emb(q, k, freqs_cos, freqs_sin) # 再做 RoPE |
| ``` |
|
|
| ### 为什么在 RoPE 前? |
|
|
| RoPE 是旋转操作,不改变向量的 norm。如果在 RoPE 后做 QK-Norm,会破坏旋转角度的相对性。 |
|
|
| > 面试点:QK-Norm 和 Pre-Norm 的区别?→ Pre-Norm 是对整个残差分支的输入做 Norm,QK-Norm 只对 Q/K 做 Norm,两者解决不同层面的稳定性问题 |
|
|
| --- |
|
|
| ## Q5. SwiGLU vs ReLU FFN(`src/core/mlp.py:7-17`) |
|
|
| ### 公式对比 |
|
|
| | | ReLU FFN | SwiGLU | |
| |---|---|---| |
| | 公式 | `max(0, xW1)W2` | `silu(xW1) ⊗ (xW2)` | |
| | 门控 | 无 | 有(W2 作为 gate) | |
| | 参数量 | 2·d·d_ff | 3·d·d_ff (实际约 2·d·d_ff) | |
| | 效果 | 一般 | 更好 | |
| |
| ### 本仓库实现(`src/core/mlp.py:7-17`) |
| |
| ```python |
| class FeedForward(nn.Module): |
| def __init__(self, config): |
| self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) |
| self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) |
| self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) |
| self.act_fn = ACT2FN[config.hidden_act] # silu |
| |
| def forward(self, x): |
| return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) |
| ``` |
| |
| ### FFN 容量公式 |
|
|
| 设 d_model=d,d_ff=f: |
| - 标准 FFN 参数:`2df` |
| - SwiGLU 参数:`3·(d·(2f/3))` ≈ `2df`(LLaMA 取 d_ff = 8d/3 配 SwiGLU) |
| |
| > 面试点:为什么 SwiGLU 比 ReLU FFN 效果好?→ 门控结构让网络学"哪些信息该通过",表达力更强、收敛更好 |
| |
| --- |
| |
| ## Q6. Pre-Norm vs Post-Norm |
| |
| ### 结构对比 |
| |
| ``` |
| Pre-Norm(本仓库): Post-Norm(原 Transformer): |
| x = x + Attn(RMSNorm(x)) x = RMSNorm(x + Attn(x)) |
| x = x + FFN(RMSNorm(x)) x = RMSNorm(x + FFN(x)) |
| ``` |
| |
| ### 为什么 Pre-Norm 更稳? |
| |
| - **Post-Norm**:梯度需要经过 Norm 层才能回传,深层梯度易消失,需 warmup |
| - **Pre-Norm**:残差通道始终直通,梯度能无损回传,训练更稳、可省/减 warmup |
| |
| ### 本仓库实现(`src/core/block.py:8-24`) |
| |
| ```python |
| class Block(nn.Module): |
| def forward(self, x, start_pos, freqs_cos, freqs_sin, mask=None): |
| # Pre-Norm 架构 |
| h = x + self.attention(self.attention_norm(x), start_pos, freqs_cos, freqs_sin, mask) |
| out = h + self.mlp(self.ffn_norm(h)) # self.mlp 可能是 FeedForward 或 MOEFeedForward |
| return out |
| ``` |
| |
| --- |
|
|
| ## Q7. Flash-Attention 为什么快? |
|
|
| ### 核心思想:IO 感知 |
|
|
| 传统注意力需要物化完整的 N×N 注意力矩阵,显存 O(N²)。Flash-Attention 通过分块计算,避免物化完整矩阵。 |
|
|
| ### 本仓库的 Flash-Attention 条件(`src/core/attention.py:28, 44`) |
|
|
| ```python |
| # 仅在以下条件使用 Flash Attention |
| if (seq_len > 1 and |
| (not self.causal or past_key_value is None) and |
| attention_mask is None): |
| # 使用 F.scaled_dot_product_attention(PyTorch 2.0+ 内置 Flash Attention) |
| ``` |
|
|
| ### 为什么需要这些条件? |
|
|
| 1. `seq_len > 1`:单 token 无需注意力 |
| 2. `not self.causal or past_key_value is None`:Flash Attention 对 causal mask 支持有限 |
| 3. `attention_mask is None`:Flash Attention 不支持自定义 mask |
|
|
| > 面试点:Flash Attention 和标准 Attention 的计算量一样吗?→ 一样,只是 IO 优化,减少 HBM 读写 |
|
|
| --- |
|
|
| ## Q8. 参数量估算(面试手算) |
|
|
| 以 LLaMA-7B 量级为例(d=4096, layers=32, heads=32, d_ff=11008, vocab=32000): |
| |
| ``` |
| embedding : V·d ≈ 32000·4096 ≈ 131M |
| per block : 2·(d² + d·d_ff) ≈ 2·(16.8M + 45.1M) ≈ 124M |
| all blocks: 32 · 124M ≈ 3.97B |
| lm_head : V·d ≈ 131M (若共享则不计) |
| ``` |
| |
| → ≈ 6.7B,与公开 7B 吻合。手算时常忽略 layernorm/bias 小头。 |
| |
| ### 本仓库默认参数(`src/models/lm/config.py`) |
| |
| | 参数 | 默认值 | |
| |------|--------| |
| | hidden_size | 768 | |
| | num_hidden_layers | 8 | |
| | num_attention_heads | 8 | |
| | num_key_value_heads | 4 | |
| | intermediate_size | ceil(768 * π / 64) * 64 ≈ 3840 | |
| | vocab_size | 6400 | |
| | max_position_embeddings | 32768 | |
| | rope_theta | 1e6 | |
|
|
| --- |
|
|
| ## Q9. 权重绑定(Weight Tying) |
|
|
| ### 本仓库实现(`src/models/lm/model.py:52, 59-60`) |
|
|
| ```python |
| class LMForCausalLM(PreTrainedModel): |
| _tied_weights_keys = ["lm_head.weight"] # 与 embed_tokens 绑定 |
| |
| def __init__(self, config): |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) |
| # lm_head.weight 与 self.model.embed_tokens.weight 共享 |
| ``` |
|
|
| ### 为什么共享权重? |
|
|
| 1. **省参数**:词表参数减半(vocab_size × hidden_size) |
| 2. **语义对齐**:embedding 空间和 logits 空间在同一空间,语义更一致 |
| 3. **小模型更稳**:减少参数量有助于正则化 |
|
|
| --- |
|
|
| ## Q10. 推理时的 KV Cache |
|
|
| ### 核心思想 |
|
|
| 缓存已算的 K/V,每步只算新 token 的 Q 与已有 K/V 做注意力,避免重算。 |
|
|
| ### 本仓库实现(`src/core/attention.py:39-42`) |
|
|
| ```python |
| def forward(self, x, start_pos, freqs_cos, freqs_sin, mask=None): |
| # KV Cache 简单拼接实现 |
| if past_key_value is not None: |
| xk = torch.cat([past_key_value[0], xk], dim=2) |
| xv = torch.cat([past_key_value[1], xv], dim=2) |
| past_key_value = (xk, xv) |
| ``` |
|
|
| ### KV Cache 的显存占用 |
|
|
| 假设: |
| - batch_size=B, seq_len=S, num_layers=L |
| - num_kv_heads=H, head_dim=D |
| - 精度=fp16(2 bytes) |
|
|
| KV Cache 显存 = `2 × B × S × L × H × D × 2 bytes` |
|
|
| > 面试点:GQA 如何减少 KV Cache?→ KV 头数从 n_heads 减到 n_kv_heads,KV Cache 减少 n_heads/n_kv_heads 倍 |
|
|
| --- |
|
|
| ## Q11. Logits to Keep 优化(`src/models/lm/model.py:65-66`) |
|
|
| ### 问题 |
|
|
| 在生成时,我们只需要最后一个 token 的 logits,但标准实现会计算所有 token 的 logits。 |
|
|
| ### 本仓库优化 |
|
|
| ```python |
| def forward(self, input_ids, ..., logits_to_keep=1): |
| # 仅计算最后 N 个 token 的 logits |
| hidden_states = hidden_states[:, -logits_to_keep:] |
| logits = self.lm_head(hidden_states) |
| ``` |
|
|
| ### 显存节省 |
|
|
| 假设 vocab_size=64000,seq_len=32768: |
| - 不优化:`32768 × 64000 × 2 bytes ≈ 4GB` |
| - 优化后:`1 × 64000 × 2 bytes ≈ 128KB` |
|
|
| --- |
|
|
| ## Q12. MoE 设计模式(`src/core/mlp.py:20-49`) |
|
|
| ### 关键设计:死 Expert 梯度保持 |
|
|
| ```python |
| class MOEFeedForward(nn.Module): |
| def forward(self, x): |
| # ... top-k 选择 ... |
| |
| # 死 expert 梯度保持技巧 |
| y[0, 0] += 0 * sum(p.sum() for p in self.experts.parameters()) |
| return y |
| ``` |
|
|
| ### 为什么需要这个技巧? |
|
|
| 在 DDP 分布式训练中,如果某个 expert 完全没被选中,它的参数就不会有梯度,导致 DDP 梯度同步死锁。`0 * sum(p)` 让这些参数仍然出现在计算图中,保持 DDP 通信闭环。 |
|
|
| --- |
|
|
| ## Q13. 延迟 Buffer 初始化(`src/models/lm/model.py:31-33`) |
|
|
| ### 问题 |
|
|
| 两个场景会导致 RoPE 的 `freqs_cos`/`freqs_sin` 被重置为未初始化内存: |
|
|
| 1. **HF `from_pretrained`**:创建模型时在 meta device 上,非持久 buffer 之后用 `torch.empty_like` 分配(含随机/NaN 值)。 |
| 2. **`torch.compile`**:编译后的重新初始化也会使 buffer 值丢失。 |
| |
| ### 本仓库解决方案 |
| |
| ```python |
| # register_buffer 使用 persistent=False 避免冗余 |
| self.register_buffer("freqs_cos", freqs_cos, persistent=False) |
| self.register_buffer("freqs_sin", freqs_sin, persistent=False) |
| |
| def forward(self, input_ids, ...): |
| h = self.dropout(self.embed_tokens(input_ids)) |
| # 延迟初始化:检查 cos(0) 是否 != 1.0(未初始化时为 0/NaN/随机值) |
| if self.freqs_cos[0, 0] != 1.0: |
| freqs_cos, freqs_sin = precompute_freqs_cis(...) |
| self.freqs_cos = freqs_cos.to(device=h.device, dtype=h.dtype) |
| self.freqs_sin = freqs_sin.to(device=h.device, dtype=h.dtype) |
| ``` |
| |
| ### 为什么有效? |
| |
| `cos(0) = 1`,所以初始化的 buffer 第一元素恒为 1.0。检查 `!= 1.0` 能捕捉所有异常值(0、NaN、随机内存),比原来的 `== 0` 更鲁棒(未初始化内存可能是任意值而非精确零)。重算时同时 `.to(dtype=...)` 确保 dtype 与模型一致。 |
| |
| --- |
| |
| ## Q14. 原子写入 Checkpoint(`src/utils/checkpoint.py:7-10`) |
| |
| ### 问题 |
| |
| 训练过程中保存 checkpoint 时,如果中途崩溃,可能导致 checkpoint 损坏。 |
| |
| ### 本仓库解决方案 |
| |
| ```python |
| def save_checkpoint(model, optimizer, scheduler, path): |
| # 先保存到临时文件 |
| tmp_path = path + ".tmp" |
| torch.save({...}, tmp_path) |
| # 原子替换 |
| os.replace(tmp_path, path) |
| ``` |
| |
| ### 为什么用 os.replace? |
| |
| `os.replace` 是原子操作,要么完成要么不发生,不会出现部分写入的情况。 |
| |
| --- |
| |
| ## Q15. SkipBatchSampler(`src/utils/training.py:177-200`) |
| |
| ### 问题 |
| |
| 分布式训练中断后,需要跳过已完成的 batch 继续训练。 |
| |
| ### 本仓库解决方案 |
| |
| ```python |
| class SkipBatchSampler: |
| def __init__(self, dataset, batch_size, step): |
| self.dataset = dataset |
| self.batch_size = batch_size |
| self.step = step # 已完成的 step 数 |
| |
| def __iter__(self): |
| # 跳过前 step 个 batch |
| indices = list(range(len(self.dataset))) |
| indices = indices[self.step * self.batch_size:] |
| # 返回剩余 batch |
| ... |
| ``` |
| |
| ### 为什么需要这个? |
| |
| DDP 训练中,每个 rank 可能中断在不同的 step,需要精确跳过已完成的 batch,避免重复训练。 |
| |