速查表
17 步一页看完,面试前过一遍用。点步骤名跳到实验台对应那一步。打开"自测"会把公式和要点模糊掉,先自己回想,再点开核对。
| # | 步骤 | 代码 | 形状 | 关键公式 | 要记住的一句话 |
|---|---|---|---|---|---|
| 1 | 全局地图 | model(input_ids) |
(B,T) → (B,T,V) | logits = lm_head ∘ ln_f ∘ block1 ∘ block0 ∘ embed(ids) | 参数总数 135168;lm_head 和 wte 共用权重 |
| 2 | input_ids | torch.randint(0, V, (B, T)) |
(2, 16) | id ∈ {0, …, V−1},对应 one-hot eid ∈ ℝV | id 超出 0 到 V−1 的范围,查表会报错 |
| 3 | wte 查词表 | tok = tf.wte(input_ids) |
(2,16) → (2,16,64) | tok = E[id] = eidTE, E ∈ ℝV×C | 查表等价于 one-hot 乘矩阵,只有用到的行有梯度 |
| 4 | wpe 查位置表 | pos = tf.wpe(torch.arange(T)) |
(16, 64) | pos[t] = P[t] | 不加位置向量时,attention 分不出词序(置换等变) |
| 5 | x = tok + pos | x = tok + pos |
(2,16,64) | (16,64) 广播成 (1,16,64) | 从最后一维往前对齐,维度相等或其中一个为 1 才能相加 |
| 6 | LayerNorm | h = block.ln_1(x) |
形状不变 | y = γ ⊙ (x − μ) / √(σ² + ε) + β | 对每个 token 的 C 个数单独做;方差除以 C,不是 C−1 |
| 7 | q、k、v | qkv = c_attn(h); q,k,v = qkv.split(C, -1) |
(2,16,192) → 3 × (2,16,64) | [q | k | v] = h [Wq | Wk | Wv] | 一次乘 3C 列 = 分三次各乘 C 列 |
| 8 | 拆成多头 | q.view(B,T,H,D).transpose(1,2) |
(2,2,16,32) | q'[b,h,t,d] = q[b,t,h·D+d] | transpose 后内存不连续,合并时用 reshape 而不是 view |
| 9 | 注意力分数 | att = q @ k.transpose(-2,-1) / math.sqrt(D) |
(2,2,16,16) | S = QKT / √D | Var(q·k) = D,除以 √D 后方差回到 1 |
| 10 | 因果 mask | att.masked_fill(~tril, float('-inf')) |
形状不变 | S'ij = Sij(j ≤ i),−∞(j > i) | exp(−∞) = 0;填 0 不行,因为 exp(0) = 1 |
| 11 | softmax | att = F.softmax(att, dim=-1) |
形状不变 | aj = esj / Σk esk | 加常数不变;∂ai/∂sj = ai(δij − aj) |
| 12 | 加权求和 + 残差 | y = (att @ v).transpose(1,2).reshape(B,T,C); x = x + c_proj(y) |
(2,2,16,32) → (2,16,64) | yi = Σj Aij vj; x ← x + f(LN(x)) | 残差的导数是 I + ∂f/∂x,梯度能原样传回前面 |
| 13 | MLP | c_proj(gelu(c_fc(ln_2(x)))) |
64 → 128 → 64 | GELU(hW1 + b1) W2 + b2 | 去掉激活函数,两层就合并成一层 W1W2 |
| 14 | 两层叠起来 | for block in tf.h: ... |
始终 (2,16,64) | x(L) = x(0) + Σl (al + ml) | 最后的 x = 初始嵌入 + 所有子层输出之和 |
| 15 | ln_f + lm_head | logits = lm_head(ln_f(x)) |
(2,16,64) → (2,16,1000) | zt = LN(xt) ET | 每个词的分数 = x 和这个词向量的点积;共用 E 省 V·C = 64000 个参数 |
| 16 | loss | F.cross_entropy(logits[:,:-1].reshape(-1,V), ids[:,1:].reshape(-1)) |
N = B(T−1) = 30 | L = −(1/N) Σn log pn,yn | 没训练时 L ≈ ln V = 6.908;∂L/∂z = (p − onehot)/N |
| 17 | 对答案 | torch.allclose(logits, ref.logits, atol=1e-5) |
— | |a − b| ≤ atol + rtol·|b| | no_grad 省内存;eval() 关掉 dropout;is 比较是否同一个对象 |