Hang Zhengyang

LLM 数学速查(配合 mini-LLM)

1. 维度约定

  • B:batch(本项目很多路径是单样本)
  • T:序列长度
  • D:d_model
  • H:num_heads
  • Hkv:num_kv_heads
  • Dh:head_dim
  • F:d_ff
  • V:词表大小

约束:

D=H·Dh

解释:模型总通道数 D 会被切成 H 个头,每个头宽度是 Dh。
如果这个式子不成立,reshape/split head 会直接维度错误,attention 无法执行。

2. 注意力核心公式

Q=XWQ,K=XWK,V=XWV

解释:

  • X:输入隐藏状态(每行一个 token)。
  • WQ,WK,WV:三组可学习投影矩阵。
  • 这一步在做“角色分工”:
    Q 负责“我要找什么”,K 负责“我有什么特征”,V 负责“我能提供什么内容”。
A=softmax(QK⊤Dh+M)

解释:

  • QK⊤:两两相似度打分,分数越大表示“越相关”。
  • 除以 Dh:防止维度大时分数过大,softmax 过饱和导致梯度变差。
  • M:mask(因果 mask 把未来位置设为极小值,确保“不能偷看未来”)。
  • softmax:把每行打分转成概率分布(每行和为 1),表示“关注权重”。
Y=AV

解释:

  • 用权重矩阵 A 对值向量 V 做加权求和。
  • 输出 Y 的每一行是“当前 token 从上下文聚合到的信息”。
  • 本质是可学习的信息路由:哪些历史 token 更重要,就给更高权重。

3. RMSNorm

RMSNorm(x)=γ⊙x1D∑ixi2+ϵ

解释:

  • 先计算向量 x 的均方根(RMS),再按 RMS 缩放,最后乘可学习系数 γ。
  • ϵ 是数值稳定项,避免除 0。
  • 与 LayerNorm 的区别:RMSNorm 不做“减均值中心化”,只做尺度归一化,通常更轻量。
  • 作用:让不同 token 的激活尺度更稳定,优化更容易收敛。

4. SwiGLU

g=XWg,u=XWu,h=SiLU(g)⊙u,y=hWd

解释:

  • 先从同一输入 X 走两条支路:门控支路 g 和内容支路 u。
  • SiLU(g) 作为连续门控系数,和 u 做逐元素乘法(⊙)。
  • 再通过 Wd 投影回模型维度。
  • 直觉:模型先决定“放多少内容通过”,再输出,通常比普通 FFN 表达力更好。
SiLU(z)=z·σ(z)

解释:

  • σ(z) 是 sigmoid,取值在 (0,1)。
  • 当 z 大时,σ(z)≈1,输出接近线性;当 z 小时会被抑制。
  • 这让激活既有非线性,又不会像硬门控那样不连续。

5. 交叉熵与梯度

softmax 概率:

pi=ezi∑jezj

解释:

  • zi 是第 i 个词的 logit 打分。
  • 指数让高分项更突出,分母归一化得到概率。
  • 这一步把“分数”变成“可比较的概率分布”。

单样本交叉熵(标签 y):

ℒ=−logpy

解释:

  • 只看正确类别 y 的概率。
  • py 越大,loss 越小;py 越小,loss 急剧增大。
  • 所以它强烈惩罚“很自信但错”的预测。

对 logits 梯度:

∂ℒ∂zi=pi−1[i=y]

解释:

  • 对正确类(i=y):梯度是 py−1(通常负),更新会推高正确类 logit。
  • 对错误类(i≠y):梯度是 pi(正),更新会压低错误类 logit。
  • 这条梯度式子直接体现了“把概率质量从错类挪到对类”。

6. AdamW

mt=β1mt−1+(1−β1)gt,vt=β2vt−1+(1−β2)gt2

解释:

  • gt:当前梯度。
  • mt:一阶动量(梯度方向的指数滑动平均),让更新更平滑。
  • vt:二阶动量(梯度平方的指数滑动平均),估计每个参数维度的尺度。
w←w−η(m^tv^t+ϵ+λw)

解释:

  • η:学习率。
  • m^tv^t+ϵ:按维度自适应调步长,降低“某些维度过大梯度”的影响。
  • +λw:decoupled weight decay,直接收缩参数,帮助正则化。
  • 整体效果:比裸 SGD 更稳,尤其在不同参数尺度差异较大的网络里。

7. KV Cache 的复杂度直觉

  • Prefill:仍接近 O(T2)(构建上下文)
  • Decode(每步):O(T) 与历史长度线性相关
  • 多步生成累计比“每步重算全历史”显著更省

解释:

  • Prefill:第一次处理整段 prompt,注意力仍是“全长对全长”,所以近似二次复杂度。
  • Decode:每次只新增 1 个 token,查询当前 Qt 对历史 K1:t,V1:t,单步是线性复杂度。
  • 为什么省:历史 K,V 不重复计算,缓存后只追加新 token 的 Kt,Vt。