omni / docs /interview /transformer-arch.md
chenbhao's picture
docs: update README and docs for model_type rename, RoPE guard fix, config paths
e14a9b6
|
Raw
History Blame Contribute Delete
15.8 kB

面试:Transformer 架构深度

本仓库 src/core/ 的纯 Transformer 组件,对应代码:norm.pyrope.pyattention.pymlp.pyblock.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)

本仓库 LMsrc/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

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

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

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

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

# 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

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

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

# 仅在以下条件使用 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

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

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。

本仓库优化

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 梯度保持

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 值丢失。

本仓库解决方案

# 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 损坏。

本仓库解决方案

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 继续训练。

本仓库解决方案

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,避免重复训练。