ckd0817/LLM-Interview-Code

Python

870

26 commits

updated Jul 26, 2026

See the code

README

LLM 面试手撕代码大全

大模型面试必备:从注意力机制到强化学习,从零实现核心组件

Python PyTorch License

目录

项目简介

本项目收录了大语言模型(LLM)面试中高频出现的手撕代码实现,涵盖:

  • 注意力机制:MHA、GQA、MLA 等现代注意力变体
  • 归一化层:LayerNorm、RMSNorm
  • 位置编码:RoPE 旋转位置编码
  • 前馈网络:FFN、SwiGLU、MoE
  • 损失函数:Pretrain、SFT、DPO、PPO、GRPO 等训练损失
  • 参数高效微调:LoRA

项目特色

  • 从零实现,无第三方依赖
  • 详细注释,张量形状图解
  • 公式推导,原理解析

写给深度学习初学者

初学深度学习时,我常感到困惑:为什么对相同的张量做同样的操作,得到的结果却不同?比如 Q、K、V 三个矩阵的形成过程,代码完全一致,却产生不同的表达。

后来我逐渐理解:张量操作的代码是相同的,但运行时权重矩阵的参数不同。这些参数的更新由反向传播自动完成,我们无法直接控制。如果想深入了解参数是如何形成的,需要学习反向传播的原理。我们能做的是学会正确操作张量,设计合理的计算图,让反向传播按照预期方向更新参数。

仔细观察这些代码,你会发现它们大多是在进行维度变化和对齐的操作——reshape、transpose、expand、concatenate……掌握这些操作,就能理解数据在网络中是如何流动的。

因此,掌握这些 LLM 组件的本质是理解并记忆各种张量操作。面试手撕代码时,考察的重点也是张量操作和原理的记忆,而非参数如何形成。张量操作是理解这一切的基础

为此,我专门准备了一个教程:PyTorch 张量变换与重塑教程


项目结构

LLM-Interview-Code/
├── attention/                     # 注意力机制
│   ├── ScaledDotProductAttention.py
│   ├── MultiHeadAttention.py
│   ├── GroupQueryAttention.py
│   └── MultiLatentAttention.py
├── normalization/                 # 归一化层
│   ├── LayerNorm.py
│   └── RMSNorm.py
├── position/                      # 位置编码
│   └── RotaryEmbedding.py
├── ffn/                           # 前馈网络
│   ├── FFN.py
│   ├── SwiGLUFFN.py
│   └── MoE.py
├── loss/                          # 损失函数
│   ├── SFTLoss.py
│   ├── DPOLoss.py
│   ├── PPOLoss.py
│   ├── GRPOLoss.py
│   ├── PretainLoss.py
│   └── EntropyLoss.py
├── peft/                          # 参数高效微调
│   └── LoRALinear.py
├── pytorch_tensor_reshape.ipynb   # PyTorch 张量操作教程
└── README.md

注意力机制

实现说明 本仓库里的注意力模块默认使用 dropout_p=0.0,更贴近近两年主流 decoder-only LLM 的常见配置。 如果你是为了讲解经典 Transformer 正则化,或者在小数据训练中想显式增加随机性,可以手动传入非零 dropout。

Scaled Dot-Product Attention

背景与动机

缩放点积注意力(Scaled Dot-Product Attention)是所有注意力机制的基础。它计算 Query 和 Key 的点积,除以缩放因子后通过 softmax 得到注意力权重,最后加权求和 Value。

这是理解多头注意力的前提,也是面试中最常考的基础版本。

核心公式

$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$

为什么需要缩放?

  • 当 $d_k$ 较大时,点积结果也会很大
  • 过大的值进入 softmax 后会导致梯度消失
  • 除以 $\sqrt{d_k}$ 使方差稳定在 1 附近

张量形状流程图

flowchart TD
    Q["Q: [batch, num_heads, seq_len, head_dim]"] --> Scores["scores = Q @ K^T / sqrt(head_dim)"]
    K["K: [batch, num_heads, seq_len, head_dim]"] --> Scores
    Scores --> ScoresShape["[batch, num_heads, seq_len, seq_len]"]
    ScoresShape --> Softmax["attn_weights = softmax(scores, dim=-1)"]
    Softmax --> WeightsShape["[batch, num_heads, seq_len, seq_len]"]
    WeightsShape --> WeightedSum["output = attn_weights @ V"]
    V["V: [batch, num_heads, seq_len, head_dim]"] --> WeightedSum
    WeightedSum --> Output["输出: [batch, num_heads, seq_len, head_dim]"]

Multi-Head Attention

背景与动机

多头注意力(Multi-Head Attention, MHA)通过将输入映射到多个子空间并行计算注意力,模型可以同时关注不同位置的不同表示子空间信息。

每个头学习不同的注意力模式,最后拼接并通过输出投影融合。

核心公式

$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h)W^O$$

其中 $\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> WQ["W_q"]
    X --> WK["W_k"]
    X --> WV["W_v"]
    WQ --> Q0["Q<br/>[batch, seq_len, model_dim]"]
    WK --> K0["K<br/>[batch, seq_len, model_dim]"]
    WV --> V0["V<br/>[batch, seq_len, model_dim]"]
    Q0 --> RQ["reshape"] --> Q["Q<br/>[batch, num_heads, seq_len, head_dim]"]
    K0 --> RK["reshape"] --> K["K<br/>[batch, num_heads, seq_len, head_dim]"]
    V0 --> RV["reshape"] --> V["V<br/>[batch, num_heads, seq_len, head_dim]"]
    Q --> Attention["Scaled Dot-Product Attention"]
    K --> Attention
    V --> Attention
    Attention --> Context["[batch, num_heads, seq_len, head_dim]"]
    Context --> Concat["reshape (concat heads)"]
    Concat --> ConcatShape["[batch, seq_len, model_dim]"]
    ConcatShape --> WO["W_o (output projection)"]
    WO --> Output["输出: [batch, seq_len, model_dim]"]

Group Query Attention

背景与动机

分组查询注意力(Grouped Query Attention, GQA)是 MHA 和 Multi-Query Attention (MQA) 的折中方案。在 GQA 中,Query 有 H 个头,而 Key 和 Value 只有 G 个头(G < H),多组 Query 共享同一组 K/V。

这显著减少了 KV Cache 的显存占用,同时保持了较好的模型质量。LLaMA 2、LLaMA 3 等模型都采用了 GQA。

对比

类型Q 头数K 头数V 头数KV Cache
MHAHHH100%
GQAHGGG/H × 100%
MQAH111/H × 100%

核心公式

与 MHA 相同,但 K、V 需要通过 repeat_kv 扩展到与 Q 相同的头数:

K_expanded = repeat(K, num_heads // num_kv_heads)
V_expanded = repeat(V, num_heads // num_kv_heads)

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> WQ["W_q"]
    X --> WK["W_k"]
    X --> WV["W_v"]
    WQ --> Q["Q<br/>[batch, num_heads, seq_len, head_dim]"]
    WK --> K["K<br/>[batch, num_kv_heads, seq_len, head_dim]"]
    WV --> V["V<br/>[batch, num_kv_heads, seq_len, head_dim]"]
    K --> RepeatK["repeat_kv"] --> KExp["K_exp<br/>[batch, num_heads, seq_len, head_dim]"]
    V --> RepeatV["repeat_kv"] --> VExp["V_exp<br/>[batch, num_heads, seq_len, head_dim]"]
    Q --> Attention["Attention"]
    KExp --> Attention
    VExp --> Attention
    Attention --> Output["输出: [batch, seq_len, model_dim]"]

Multi-Latent Attention

实现说明 MultiLatentAttention.py 是用于讲解核心张量变换的简化版 MLA:保留 Q/KV 低秩投影与 RoPE 注意力主线,但未实现原论文中的共享解耦 RoPE Key、 投影吸收和增量 KV Cache 接口。

背景与动机

多潜变量注意力(Multi-Latent Attention, MLA)由 DeepSeek-V2 提出,通过将 KV 压缩到低维潜空间来大幅减少 KV Cache。与 GQA 不同,MLA 不是通过减少头数,而是通过降维压缩来实现内存节省。

核心思想

  • KV 先通过下投影压缩到潜空间(存入 Cache)
  • 计算注意力时再上投影恢复
  • 压缩比可达 90%+,同时保持性能

核心公式

KV 压缩(下投影到潜空间):

$$c_{KV} = W_{DKV} \cdot h_t$$

KV 恢复(上投影恢复 K、V):

$$[k_{t}, v_{t}] = W_{UKV} \cdot c_{KV}$$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> KVPath["KV路径"]
    X --> QPath["Q路径"]
    KVPath --> KVDown["kv_down"]
    KVDown --> KVLatent["[batch, seq_len, latent_dim]<br/>存入 KV Cache"]
    KVLatent --> KVUp["kv_up"]
    KVUp --> KV["[K, V]"]
    QPath --> QDown["q_down"]
    QDown --> QLatent["[batch, seq_len, latent_dim]<br/>仅用于当前 Query"]
    QLatent --> QUp["q_up"]
    QUp --> Q["[Q, Q_rope]"]
    KV --> Attention["RoPE + Attention"]
    Q --> Attention
    Attention --> Output["输出: [batch, seq_len, model_dim]"]

归一化层

LayerNorm

背景与动机

层归一化(Layer Normalization)在每个样本的特征维度上进行归一化,使得训练更加稳定。与 BatchNorm 不同,LayerNorm 不依赖 batch size,因此更适合序列模型和 Transformer。

核心公式

$$\text{LN}(x) = \gamma \cdot \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta$$

其中:

  • $\mu = \frac{1}{d}\sum_{i=1}^{d} x_i$ (均值)
  • $\sigma^2 = \frac{1}{d}\sum_{i=1}^{d} (x_i - \mu)^2$ (方差)
  • $\gamma, \beta$:可学习的缩放和偏移参数

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> Mean["mean(x, dim=-1)"]
    Mean --> MeanShape["[batch, seq_len, 1]"]
    MeanShape --> Var["var(x, dim=-1)"]
    Var --> VarShape["[batch, seq_len, 1]"]
    VarShape --> Normalize["(x - mean) / sqrt(var + eps)"]
    X --> Normalize
    Normalize --> NormShape["[batch, seq_len, model_dim]"]
    NormShape --> Affine["x * gamma + beta"]
    Affine --> Output["输出: [batch, seq_len, model_dim]"]

RMSNorm

背景与动机

均方根归一化(Root Mean Square Normalization)是 LayerNorm 的简化版本。它移除了均值计算,只使用 RMS 进行归一化。这种方法计算更快,且在很多 LLM(如 LLaMA、Mistral)中表现优异。

核心公式

$$\text{RMSNorm}(x) = \gamma \cdot \frac{x}{\sqrt{\frac{1}{d}\sum_{i=1}^{d} x_i^2 + \epsilon}}$$

与 LayerNorm 的区别

  • 不计算均值(去中心化)
  • 只有一个可学习参数 $\gamma$(无 $\beta$)
  • 计算量更少,推理更快

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> RMS["rms = sqrt(mean(x^2, dim=-1) + eps)<br/>[batch, seq_len, 1]"]
    RMS --> Normalize["x / rms<br/>[batch, seq_len, model_dim]"]
    X --> Normalize
    Normalize --> Scale["x * gamma"]
    Scale --> Output["输出: [batch, seq_len, model_dim]"]

位置编码

Rotary Position Embedding (RoPE)

背景与动机

旋转位置编码(RoPE)通过旋转向量的方式注入位置信息,具有以下优势:

  • 相对位置感知:自动捕捉 token 之间的相对位置关系
  • 外推能力:可以处理比训练时更长的序列
  • 计算高效:通过逐元素乘法实现

目前被 LLaMA、Mistral、Qwen 等主流模型采用。

核心公式

对于位置 $m$ 的向量 $x$,RoPE 将其旋转:

\begin{bmatrix}
x_1' \\
x_2'
\end{bmatrix}
=
\begin{bmatrix}
\cos(m\theta) & -\sin(m\theta) \\
\sin(m\theta) & \cos(m\theta)
\end{bmatrix}
\cdot
\begin{bmatrix}
x_1 \\
x_2
\end{bmatrix}

展开形式:

$$x_1' = x_1 \cos(m\theta) - x_2 \sin(m\theta)$$

$$x_2' = x_1 \sin(m\theta) + x_2 \cos(m\theta)$$

其中 $\theta_i = 10000^{-2i/d}$

实现原理

位置 m 的旋转角度: θ_m = m * θ_base
预计算 cos(m*θ) 和 sin(m*θ) 用于所有位置

旋转公式:
[x1', x2'] = [x1*cos - x2*sin, x1*sin + x2*cos]

等价于:
x' = x * cos + rotate_half(x) * sin
其中 rotate_half(x) = [-x后半, x前半]

张量形状流程图

flowchart TD
    QK["输入 Q, K: [batch, seq_len, num_heads, head_dim]"]
    Precompute["预计算 cos, sin: [max_seq_len, head_dim]"]
    Precompute --> Slice["取当前序列长度: cos[:seq_len], sin[:seq_len]<br/>[1, seq_len, 1, head_dim]"]
    QK --> Rotate["rotate_half(Q) = [-Q后半, Q前半]"]
    Slice --> Apply["Q_rotated = Q * cos + rotate_half(Q) * sin<br/>K_rotated = K * cos + rotate_half(K) * sin"]
    Rotate --> Apply
    QK --> Apply
    Apply --> Output["输出: [batch, seq_len, num_heads, head_dim]"]

前馈网络

FFN

背景与动机

前馈网络(Feed-Forward Network)是 Transformer 中注意力层之后的两层全连接网络,用于对特征进行非线性变换。它是 Transformer 中参数量最大的部分。

核心公式

$$\text{FFN}(x) = W_2 \cdot \text{ReLU}(W_1 x)$$

通常 $d_{ff} = 4 \times d_{model}$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> W1["W_1 (up projection)"]
    W1 --> Hidden["[batch, seq_len, 4*model_dim]"]
    Hidden --> ReLU["ReLU"]
    ReLU --> W2["W_2 (down projection)"]
    W2 --> Output["输出: [batch, seq_len, model_dim]"]

SwiGLU

背景与动机

SwiGLU 是 GLU(Gated Linear Unit)变体之一,被 LLaMA、PaLM 等模型采用。相比标准 FFN,SwiGLU 引入门控机制和 Swish 激活函数,提升了模型性能。

核心公式

$$\text{SwiGLU}(x) = (W_{gate}(x) \odot \text{SiLU}(W_{up}(x))) \cdot W_{down}$$

其中:

  • $\text{SiLU}(x) = x \cdot \sigma(x)$ (也称为 Swish)
  • $\odot$ 表示逐元素乘法

参数量对比:标准 FFN 有 2 个矩阵,SwiGLU 有 3 个矩阵

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> WGate["W_gate"]
    X --> WUp["W_up"]
    WGate --> Gate["gate<br/>[batch, seq_len, intermediate_dim]"]
    Gate --> SiLU["SiLU"]
    WUp --> Up["up<br/>[batch, seq_len, intermediate_dim]"]
    SiLU --> Multiply["gate * up<br/>门控乘法"]
    Up --> Multiply
    Multiply --> Hidden["[batch, seq_len, intermediate_dim]"]
    Hidden --> WDown["W_down"]
    WDown --> Output["输出: [batch, seq_len, model_dim]"]

Mixture of Experts

背景与动机

混合专家模型(Mixture of Experts, MoE)通过稀疏激活实现模型容量的极大扩展。每个 token 只激活部分专家网络,使得总参数量可以很大,但计算量保持可控。

核心思想

  • Router 决定每个 token 应该由哪些专家处理
  • Top-K 路由:每个 token 只激活 K 个专家
  • 专家输出按路由权重加权求和

代表模型:Mixtral 8x7B、DeepSeek-V2、GPT-4 等。

核心公式

$$\text{MoE}(x) = \sum_{i \in \text{TopK}} \text{softmax}(\text{router}(x))_i \cdot E_i(x)$$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> Flatten["flatten: [batch*seq_len, model_dim]"]
    Flatten --> Router["Router"]
    Flatten --> Experts["Experts (E_1, ..., E_N)"]
    Router --> Logits["[batch*seq_len, num_experts]"]
    Logits --> TopK["Top-K"]
    TopK --> Softmax["softmax"]
    Softmax --> Mask["mask"]
    Experts --> ExpertOutput["expert_output"]
    Mask --> WeightedSum["weighted sum (按路由权重累加)"]
    ExpertOutput --> WeightedSum
    WeightedSum --> Output["输出 reshape: [batch, seq_len, model_dim]"]

损失函数

Pretrain Loss

背景与动机

预训练损失(Pretrain Loss)是因果语言模型最基础的训练目标:给定当前位置之前的 token,预测下一个 token。除 padding 等需要忽略的位置外,序列中的所有 token 都参与损失计算。

这是所有 LLM 训练的基础,理解它是学习 SFT、DPO、PPO 等训练方法的前提。

核心公式

$$\mathcal{L}{\text{Pretrain}} = -\sum{t=2}^{T} \log P(x_t \mid x_{<t})$$

张量形状流程图

flowchart TD
    Logits["logits: [batch, seq_len, vocab_size]"] --> Shift
    Labels["labels: [batch, seq_len]"] --> Shift
    Shift["Shift (自回归预测)<br/>logits[:, :-1, :]<br/>labels[:, 1:]"]
    Shift --> Flatten["Flatten<br/>[batch*(seq_len-1), vocab_size]<br/>[batch*(seq_len-1)]"]
    Flatten --> Loss["CrossEntropy<br/>ignore_index = -100"]
    Loss --> Output["loss: scalar"]

SFT Loss

背景与动机

监督微调(Supervised Fine-Tuning, SFT)损失是带 prompt 掩码的交叉熵损失。在指令微调中,通常只计算 response 部分的损失,不计算 prompt 部分。

如果只看损失函数的实现,SFT Loss 与 Pretrain Loss 的 next-token Cross Entropy 完全相同,唯一的实质差异是 mask

Pretrain LossSFT Loss
参与损失的 token所有未被忽略的 token仅 response token
prompt label正常参与损失设为 -100,不参与损失
Shift + CrossEntropy相同相同

两者的训练阶段和数据形式不同,但损失计算的主干没有变化。

核心公式

$$\mathcal{L}{SFT} = -\sum{t=p}^{T} \log P(y_t | x, y_{<t})$$

其中 $p$ 是 prompt 长度,即只对 response 部分计算损失。

张量形状流程图

flowchart TD
    Logits["logits: [batch, seq_len, vocab_size]"] --> Mask
    Labels["labels: [batch, seq_len]"] --> Mask
    PromptLengths["prompt_lengths: [batch]"] --> Mask
    Mask["构造 masked_labels<br/>prompt 部分设为 -100"]
    Mask --> Shift["Shift (自回归预测)<br/>logits[:, :-1, :]<br/>masked_labels[:, 1:]"]
    Shift --> Loss["Flatten + CrossEntropy<br/>ignore_index = -100"]
    Loss --> Output["loss: scalar"]

DPO Loss

背景与动机

直接偏好优化(Direct Preference Optimization, DPO)是一种无需奖励模型的 RLHF 替代方案。它直接在偏好数据上优化策略,简化了训练流程。

核心思想:增加 chosen 回答的概率,降低 rejected 回答的概率

核心公式

$$\mathcal{L}{DPO} = -\mathbb{E}\left[\log \sigma\left(\beta \left(\log \frac{\pi\theta(y_w|x)}{\pi_{ref}(y_w|x)} - \log \frac{\pi_\theta(y_l|x)}{\pi_{ref}(y_l|x)}\right)\right)\right]$$

其中:

  • $y_w$:chosen(优选)回答
  • $y_l$:rejected(拒绝)回答
  • $\pi_\theta$:当前策略
  • $\pi_{ref}$:参考策略(通常是 SFT 模型)
  • $\beta$:KL 散度约束系数

PPO Loss

背景与动机

近端策略优化(Proximal Policy Optimization, PPO)通过裁剪重要性采样比率来限制策略更新幅度,防止策略崩溃。是 RLHF 训练的核心算法。

核心公式

$$\mathcal{L}_{PPO} = \mathbb{E}\left[\min\left(r_t(\theta) \hat{A}_t, \text{clip}(r_t, 1-\epsilon, 1+\epsilon)\hat{A}_t\right)\right]$$

其中:

  • $r_t = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{old}}(a_t|s_t)}$(重要性采样比率)
  • $\hat{A}_t$:优势函数估计
  • $\epsilon$:裁剪参数(通常 0.2)

裁剪机制图解

flowchart LR
    Advantage{"优势函数 A"}
    Advantage --> Positive["当 A > 0 (好动作)"]
    Advantage --> Negative["当 A < 0 (坏动作)"]
    Positive --> PositiveNormal["r 在 [1-ε, 1+ε] 内正常更新"]
    Positive --> PositiveClip["r > 1+ε 时停止奖励<br/>防止过度优化"]
    Negative --> NegativeNormal["r 在 [1-ε, 1+ε] 内正常惩罚"]
    Negative --> NegativeClip["r < 1-ε 时停止惩罚<br/>防止过度惩罚"]

GRPO Loss

背景与动机

分组相对策略优化(Group Relative Policy Optimization, GRPO)是 DeepSeek-Math 提出的算法,被 DeepSeek-R1 用于强化学习训练。它通过组内相对优势来消除对价值网络(Critic)的依赖。

核心思想

  • 对同一问题生成多个回答(组)
  • 在组内计算相对优势(而非绝对优势)
  • 无需训练 Critic 网络,简化训练流程

核心公式

1. 组内相对优势(Group-Relative Advantage)

对同一问题生成 $G$ 个回答,计算每个回答的奖励 $r_1, r_2, ..., r_G$,然后计算组内标准化优势:

$$\hat{A}_i = \frac{r_i - \text{mean}(\mathbf{r})}{\text{std}(\mathbf{r})}$$

2. GRPO 损失函数

$$\mathcal{L}{GRPO} = -\mathbb{E}\left[\frac{1}{G}\sum{i=1}^{G} \min\left(\rho_i \hat{A}_i, \text{clip}(\rho_i, 1-\epsilon, 1+\epsilon)\hat{A}i\right) - \beta \cdot \mathbb{D}{KL}\right]$$

其中:

  • $\rho_i = \frac{\pi_\theta(o_i|q)}{\pi_{\theta_{old}}(o_i|q)}$(重要性采样比率)
  • $\hat{A}_i$:组内相对优势
  • $\beta$:KL 散度惩罚系数
  • $\mathbb{D}_{KL}$:策略与参考策略的 KL 散度

与 PPO 的区别

特性PPOGRPO
优势估计需要 Critic 网络组内相对优势
额外网络需要 Value Head不需要
内存占用较高较低
适用场景通用 RL多候选生成场景

参数高效微调

LoRA

背景与动机

低秩适应(Low-Rank Adaptation, LoRA)通过在预训练权重旁添加低秩分解矩阵来实现参数高效微调。它冻结预训练权重,只训练少量参数,大大降低了微调成本。

核心思想:权重更新 $\Delta W$ 可以被低秩分解为 $B \cdot A$

核心公式

$$h = W_0 x + \Delta W x = W_0 x + BAx$$

其中:

  • $W_0 \in \mathbb{R}^{d \times k}$:冻结的预训练权重
  • $A \in \mathbb{R}^{r \times k}$:可训练,使用随机初始化
  • $B \in \mathbb{R}^{d \times r}$:可训练,初始化为零
  • $r \ll \min(d, k)$:低秩维度

关键设计:$B$ 初始化为零,使得初始状态 $BA = 0$,保证微调开始时模型行为不变。

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, in_features]"]
    X --> W0["W_0<br/>(frozen)"]
    X --> LoRA["LoRA分支"]
    W0 --> Base["W_0@x"]
    LoRA --> Dropout["Dropout"]
    Dropout --> A["lora_A: [rank, in_features]"]
    A --> B["lora_B: [out_features, rank]"]
    B --> Scaling["* scaling (α/r)"]
    Scaling --> Delta["B@A@x"]
    Base --> Add["+"]
    Delta --> Add
    Add --> Output["输出: [batch, seq_len, out_features]"]

参考文献

注意力机制

位置编码

归一化

前馈网络

训练方法

参数高效微调


Contributors

ckd0817

25 commits

chxu2000

1 commits

ckd0817/LLM-Interview-Code

Python

870

26 commits

updated Jul 26, 2026

See the code

README

LLM 面试手撕代码大全

大模型面试必备:从注意力机制到强化学习,从零实现核心组件

Python PyTorch License

目录

项目简介

本项目收录了大语言模型(LLM)面试中高频出现的手撕代码实现,涵盖:

  • 注意力机制:MHA、GQA、MLA 等现代注意力变体
  • 归一化层:LayerNorm、RMSNorm
  • 位置编码:RoPE 旋转位置编码
  • 前馈网络:FFN、SwiGLU、MoE
  • 损失函数:Pretrain、SFT、DPO、PPO、GRPO 等训练损失
  • 参数高效微调:LoRA

项目特色

  • 从零实现,无第三方依赖
  • 详细注释,张量形状图解
  • 公式推导,原理解析

写给深度学习初学者

初学深度学习时,我常感到困惑:为什么对相同的张量做同样的操作,得到的结果却不同?比如 Q、K、V 三个矩阵的形成过程,代码完全一致,却产生不同的表达。

后来我逐渐理解:张量操作的代码是相同的,但运行时权重矩阵的参数不同。这些参数的更新由反向传播自动完成,我们无法直接控制。如果想深入了解参数是如何形成的,需要学习反向传播的原理。我们能做的是学会正确操作张量,设计合理的计算图,让反向传播按照预期方向更新参数。

仔细观察这些代码,你会发现它们大多是在进行维度变化和对齐的操作——reshape、transpose、expand、concatenate……掌握这些操作,就能理解数据在网络中是如何流动的。

因此,掌握这些 LLM 组件的本质是理解并记忆各种张量操作。面试手撕代码时,考察的重点也是张量操作和原理的记忆,而非参数如何形成。张量操作是理解这一切的基础

为此,我专门准备了一个教程:PyTorch 张量变换与重塑教程


项目结构

LLM-Interview-Code/
├── attention/                     # 注意力机制
│   ├── ScaledDotProductAttention.py
│   ├── MultiHeadAttention.py
│   ├── GroupQueryAttention.py
│   └── MultiLatentAttention.py
├── normalization/                 # 归一化层
│   ├── LayerNorm.py
│   └── RMSNorm.py
├── position/                      # 位置编码
│   └── RotaryEmbedding.py
├── ffn/                           # 前馈网络
│   ├── FFN.py
│   ├── SwiGLUFFN.py
│   └── MoE.py
├── loss/                          # 损失函数
│   ├── SFTLoss.py
│   ├── DPOLoss.py
│   ├── PPOLoss.py
│   ├── GRPOLoss.py
│   ├── PretainLoss.py
│   └── EntropyLoss.py
├── peft/                          # 参数高效微调
│   └── LoRALinear.py
├── pytorch_tensor_reshape.ipynb   # PyTorch 张量操作教程
└── README.md

注意力机制

实现说明 本仓库里的注意力模块默认使用 dropout_p=0.0,更贴近近两年主流 decoder-only LLM 的常见配置。 如果你是为了讲解经典 Transformer 正则化,或者在小数据训练中想显式增加随机性,可以手动传入非零 dropout。

Scaled Dot-Product Attention

背景与动机

缩放点积注意力(Scaled Dot-Product Attention)是所有注意力机制的基础。它计算 Query 和 Key 的点积,除以缩放因子后通过 softmax 得到注意力权重,最后加权求和 Value。

这是理解多头注意力的前提,也是面试中最常考的基础版本。

核心公式

$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$

为什么需要缩放?

  • 当 $d_k$ 较大时,点积结果也会很大
  • 过大的值进入 softmax 后会导致梯度消失
  • 除以 $\sqrt{d_k}$ 使方差稳定在 1 附近

张量形状流程图

flowchart TD
    Q["Q: [batch, num_heads, seq_len, head_dim]"] --> Scores["scores = Q @ K^T / sqrt(head_dim)"]
    K["K: [batch, num_heads, seq_len, head_dim]"] --> Scores
    Scores --> ScoresShape["[batch, num_heads, seq_len, seq_len]"]
    ScoresShape --> Softmax["attn_weights = softmax(scores, dim=-1)"]
    Softmax --> WeightsShape["[batch, num_heads, seq_len, seq_len]"]
    WeightsShape --> WeightedSum["output = attn_weights @ V"]
    V["V: [batch, num_heads, seq_len, head_dim]"] --> WeightedSum
    WeightedSum --> Output["输出: [batch, num_heads, seq_len, head_dim]"]

Multi-Head Attention

背景与动机

多头注意力(Multi-Head Attention, MHA)通过将输入映射到多个子空间并行计算注意力,模型可以同时关注不同位置的不同表示子空间信息。

每个头学习不同的注意力模式,最后拼接并通过输出投影融合。

核心公式

$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h)W^O$$

其中 $\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> WQ["W_q"]
    X --> WK["W_k"]
    X --> WV["W_v"]
    WQ --> Q0["Q<br/>[batch, seq_len, model_dim]"]
    WK --> K0["K<br/>[batch, seq_len, model_dim]"]
    WV --> V0["V<br/>[batch, seq_len, model_dim]"]
    Q0 --> RQ["reshape"] --> Q["Q<br/>[batch, num_heads, seq_len, head_dim]"]
    K0 --> RK["reshape"] --> K["K<br/>[batch, num_heads, seq_len, head_dim]"]
    V0 --> RV["reshape"] --> V["V<br/>[batch, num_heads, seq_len, head_dim]"]
    Q --> Attention["Scaled Dot-Product Attention"]
    K --> Attention
    V --> Attention
    Attention --> Context["[batch, num_heads, seq_len, head_dim]"]
    Context --> Concat["reshape (concat heads)"]
    Concat --> ConcatShape["[batch, seq_len, model_dim]"]
    ConcatShape --> WO["W_o (output projection)"]
    WO --> Output["输出: [batch, seq_len, model_dim]"]

Group Query Attention

背景与动机

分组查询注意力(Grouped Query Attention, GQA)是 MHA 和 Multi-Query Attention (MQA) 的折中方案。在 GQA 中,Query 有 H 个头,而 Key 和 Value 只有 G 个头(G < H),多组 Query 共享同一组 K/V。

这显著减少了 KV Cache 的显存占用,同时保持了较好的模型质量。LLaMA 2、LLaMA 3 等模型都采用了 GQA。

对比

类型Q 头数K 头数V 头数KV Cache
MHAHHH100%
GQAHGGG/H × 100%
MQAH111/H × 100%

核心公式

与 MHA 相同,但 K、V 需要通过 repeat_kv 扩展到与 Q 相同的头数:

K_expanded = repeat(K, num_heads // num_kv_heads)
V_expanded = repeat(V, num_heads // num_kv_heads)

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> WQ["W_q"]
    X --> WK["W_k"]
    X --> WV["W_v"]
    WQ --> Q["Q<br/>[batch, num_heads, seq_len, head_dim]"]
    WK --> K["K<br/>[batch, num_kv_heads, seq_len, head_dim]"]
    WV --> V["V<br/>[batch, num_kv_heads, seq_len, head_dim]"]
    K --> RepeatK["repeat_kv"] --> KExp["K_exp<br/>[batch, num_heads, seq_len, head_dim]"]
    V --> RepeatV["repeat_kv"] --> VExp["V_exp<br/>[batch, num_heads, seq_len, head_dim]"]
    Q --> Attention["Attention"]
    KExp --> Attention
    VExp --> Attention
    Attention --> Output["输出: [batch, seq_len, model_dim]"]

Multi-Latent Attention

实现说明 MultiLatentAttention.py 是用于讲解核心张量变换的简化版 MLA:保留 Q/KV 低秩投影与 RoPE 注意力主线,但未实现原论文中的共享解耦 RoPE Key、 投影吸收和增量 KV Cache 接口。

背景与动机

多潜变量注意力(Multi-Latent Attention, MLA)由 DeepSeek-V2 提出,通过将 KV 压缩到低维潜空间来大幅减少 KV Cache。与 GQA 不同,MLA 不是通过减少头数,而是通过降维压缩来实现内存节省。

核心思想

  • KV 先通过下投影压缩到潜空间(存入 Cache)
  • 计算注意力时再上投影恢复
  • 压缩比可达 90%+,同时保持性能

核心公式

KV 压缩(下投影到潜空间):

$$c_{KV} = W_{DKV} \cdot h_t$$

KV 恢复(上投影恢复 K、V):

$$[k_{t}, v_{t}] = W_{UKV} \cdot c_{KV}$$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> KVPath["KV路径"]
    X --> QPath["Q路径"]
    KVPath --> KVDown["kv_down"]
    KVDown --> KVLatent["[batch, seq_len, latent_dim]<br/>存入 KV Cache"]
    KVLatent --> KVUp["kv_up"]
    KVUp --> KV["[K, V]"]
    QPath --> QDown["q_down"]
    QDown --> QLatent["[batch, seq_len, latent_dim]<br/>仅用于当前 Query"]
    QLatent --> QUp["q_up"]
    QUp --> Q["[Q, Q_rope]"]
    KV --> Attention["RoPE + Attention"]
    Q --> Attention
    Attention --> Output["输出: [batch, seq_len, model_dim]"]

归一化层

LayerNorm

背景与动机

层归一化(Layer Normalization)在每个样本的特征维度上进行归一化,使得训练更加稳定。与 BatchNorm 不同,LayerNorm 不依赖 batch size,因此更适合序列模型和 Transformer。

核心公式

$$\text{LN}(x) = \gamma \cdot \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta$$

其中:

  • $\mu = \frac{1}{d}\sum_{i=1}^{d} x_i$ (均值)
  • $\sigma^2 = \frac{1}{d}\sum_{i=1}^{d} (x_i - \mu)^2$ (方差)
  • $\gamma, \beta$:可学习的缩放和偏移参数

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> Mean["mean(x, dim=-1)"]
    Mean --> MeanShape["[batch, seq_len, 1]"]
    MeanShape --> Var["var(x, dim=-1)"]
    Var --> VarShape["[batch, seq_len, 1]"]
    VarShape --> Normalize["(x - mean) / sqrt(var + eps)"]
    X --> Normalize
    Normalize --> NormShape["[batch, seq_len, model_dim]"]
    NormShape --> Affine["x * gamma + beta"]
    Affine --> Output["输出: [batch, seq_len, model_dim]"]

RMSNorm

背景与动机

均方根归一化(Root Mean Square Normalization)是 LayerNorm 的简化版本。它移除了均值计算,只使用 RMS 进行归一化。这种方法计算更快,且在很多 LLM(如 LLaMA、Mistral)中表现优异。

核心公式

$$\text{RMSNorm}(x) = \gamma \cdot \frac{x}{\sqrt{\frac{1}{d}\sum_{i=1}^{d} x_i^2 + \epsilon}}$$

与 LayerNorm 的区别

  • 不计算均值(去中心化)
  • 只有一个可学习参数 $\gamma$(无 $\beta$)
  • 计算量更少,推理更快

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> RMS["rms = sqrt(mean(x^2, dim=-1) + eps)<br/>[batch, seq_len, 1]"]
    RMS --> Normalize["x / rms<br/>[batch, seq_len, model_dim]"]
    X --> Normalize
    Normalize --> Scale["x * gamma"]
    Scale --> Output["输出: [batch, seq_len, model_dim]"]

位置编码

Rotary Position Embedding (RoPE)

背景与动机

旋转位置编码(RoPE)通过旋转向量的方式注入位置信息,具有以下优势:

  • 相对位置感知:自动捕捉 token 之间的相对位置关系
  • 外推能力:可以处理比训练时更长的序列
  • 计算高效:通过逐元素乘法实现

目前被 LLaMA、Mistral、Qwen 等主流模型采用。

核心公式

对于位置 $m$ 的向量 $x$,RoPE 将其旋转:

\begin{bmatrix}
x_1' \\
x_2'
\end{bmatrix}
=
\begin{bmatrix}
\cos(m\theta) & -\sin(m\theta) \\
\sin(m\theta) & \cos(m\theta)
\end{bmatrix}
\cdot
\begin{bmatrix}
x_1 \\
x_2
\end{bmatrix}

展开形式:

$$x_1' = x_1 \cos(m\theta) - x_2 \sin(m\theta)$$

$$x_2' = x_1 \sin(m\theta) + x_2 \cos(m\theta)$$

其中 $\theta_i = 10000^{-2i/d}$

实现原理

位置 m 的旋转角度: θ_m = m * θ_base
预计算 cos(m*θ) 和 sin(m*θ) 用于所有位置

旋转公式:
[x1', x2'] = [x1*cos - x2*sin, x1*sin + x2*cos]

等价于:
x' = x * cos + rotate_half(x) * sin
其中 rotate_half(x) = [-x后半, x前半]

张量形状流程图

flowchart TD
    QK["输入 Q, K: [batch, seq_len, num_heads, head_dim]"]
    Precompute["预计算 cos, sin: [max_seq_len, head_dim]"]
    Precompute --> Slice["取当前序列长度: cos[:seq_len], sin[:seq_len]<br/>[1, seq_len, 1, head_dim]"]
    QK --> Rotate["rotate_half(Q) = [-Q后半, Q前半]"]
    Slice --> Apply["Q_rotated = Q * cos + rotate_half(Q) * sin<br/>K_rotated = K * cos + rotate_half(K) * sin"]
    Rotate --> Apply
    QK --> Apply
    Apply --> Output["输出: [batch, seq_len, num_heads, head_dim]"]

前馈网络

FFN

背景与动机

前馈网络(Feed-Forward Network)是 Transformer 中注意力层之后的两层全连接网络,用于对特征进行非线性变换。它是 Transformer 中参数量最大的部分。

核心公式

$$\text{FFN}(x) = W_2 \cdot \text{ReLU}(W_1 x)$$

通常 $d_{ff} = 4 \times d_{model}$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> W1["W_1 (up projection)"]
    W1 --> Hidden["[batch, seq_len, 4*model_dim]"]
    Hidden --> ReLU["ReLU"]
    ReLU --> W2["W_2 (down projection)"]
    W2 --> Output["输出: [batch, seq_len, model_dim]"]

SwiGLU

背景与动机

SwiGLU 是 GLU(Gated Linear Unit)变体之一,被 LLaMA、PaLM 等模型采用。相比标准 FFN,SwiGLU 引入门控机制和 Swish 激活函数,提升了模型性能。

核心公式

$$\text{SwiGLU}(x) = (W_{gate}(x) \odot \text{SiLU}(W_{up}(x))) \cdot W_{down}$$

其中:

  • $\text{SiLU}(x) = x \cdot \sigma(x)$ (也称为 Swish)
  • $\odot$ 表示逐元素乘法

参数量对比:标准 FFN 有 2 个矩阵,SwiGLU 有 3 个矩阵

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> WGate["W_gate"]
    X --> WUp["W_up"]
    WGate --> Gate["gate<br/>[batch, seq_len, intermediate_dim]"]
    Gate --> SiLU["SiLU"]
    WUp --> Up["up<br/>[batch, seq_len, intermediate_dim]"]
    SiLU --> Multiply["gate * up<br/>门控乘法"]
    Up --> Multiply
    Multiply --> Hidden["[batch, seq_len, intermediate_dim]"]
    Hidden --> WDown["W_down"]
    WDown --> Output["输出: [batch, seq_len, model_dim]"]

Mixture of Experts

背景与动机

混合专家模型(Mixture of Experts, MoE)通过稀疏激活实现模型容量的极大扩展。每个 token 只激活部分专家网络,使得总参数量可以很大,但计算量保持可控。

核心思想

  • Router 决定每个 token 应该由哪些专家处理
  • Top-K 路由:每个 token 只激活 K 个专家
  • 专家输出按路由权重加权求和

代表模型:Mixtral 8x7B、DeepSeek-V2、GPT-4 等。

核心公式

$$\text{MoE}(x) = \sum_{i \in \text{TopK}} \text{softmax}(\text{router}(x))_i \cdot E_i(x)$$

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, model_dim]"]
    X --> Flatten["flatten: [batch*seq_len, model_dim]"]
    Flatten --> Router["Router"]
    Flatten --> Experts["Experts (E_1, ..., E_N)"]
    Router --> Logits["[batch*seq_len, num_experts]"]
    Logits --> TopK["Top-K"]
    TopK --> Softmax["softmax"]
    Softmax --> Mask["mask"]
    Experts --> ExpertOutput["expert_output"]
    Mask --> WeightedSum["weighted sum (按路由权重累加)"]
    ExpertOutput --> WeightedSum
    WeightedSum --> Output["输出 reshape: [batch, seq_len, model_dim]"]

损失函数

Pretrain Loss

背景与动机

预训练损失(Pretrain Loss)是因果语言模型最基础的训练目标:给定当前位置之前的 token,预测下一个 token。除 padding 等需要忽略的位置外,序列中的所有 token 都参与损失计算。

这是所有 LLM 训练的基础,理解它是学习 SFT、DPO、PPO 等训练方法的前提。

核心公式

$$\mathcal{L}{\text{Pretrain}} = -\sum{t=2}^{T} \log P(x_t \mid x_{<t})$$

张量形状流程图

flowchart TD
    Logits["logits: [batch, seq_len, vocab_size]"] --> Shift
    Labels["labels: [batch, seq_len]"] --> Shift
    Shift["Shift (自回归预测)<br/>logits[:, :-1, :]<br/>labels[:, 1:]"]
    Shift --> Flatten["Flatten<br/>[batch*(seq_len-1), vocab_size]<br/>[batch*(seq_len-1)]"]
    Flatten --> Loss["CrossEntropy<br/>ignore_index = -100"]
    Loss --> Output["loss: scalar"]

SFT Loss

背景与动机

监督微调(Supervised Fine-Tuning, SFT)损失是带 prompt 掩码的交叉熵损失。在指令微调中,通常只计算 response 部分的损失,不计算 prompt 部分。

如果只看损失函数的实现,SFT Loss 与 Pretrain Loss 的 next-token Cross Entropy 完全相同,唯一的实质差异是 mask

Pretrain LossSFT Loss
参与损失的 token所有未被忽略的 token仅 response token
prompt label正常参与损失设为 -100,不参与损失
Shift + CrossEntropy相同相同

两者的训练阶段和数据形式不同,但损失计算的主干没有变化。

核心公式

$$\mathcal{L}{SFT} = -\sum{t=p}^{T} \log P(y_t | x, y_{<t})$$

其中 $p$ 是 prompt 长度,即只对 response 部分计算损失。

张量形状流程图

flowchart TD
    Logits["logits: [batch, seq_len, vocab_size]"] --> Mask
    Labels["labels: [batch, seq_len]"] --> Mask
    PromptLengths["prompt_lengths: [batch]"] --> Mask
    Mask["构造 masked_labels<br/>prompt 部分设为 -100"]
    Mask --> Shift["Shift (自回归预测)<br/>logits[:, :-1, :]<br/>masked_labels[:, 1:]"]
    Shift --> Loss["Flatten + CrossEntropy<br/>ignore_index = -100"]
    Loss --> Output["loss: scalar"]

DPO Loss

背景与动机

直接偏好优化(Direct Preference Optimization, DPO)是一种无需奖励模型的 RLHF 替代方案。它直接在偏好数据上优化策略,简化了训练流程。

核心思想:增加 chosen 回答的概率,降低 rejected 回答的概率

核心公式

$$\mathcal{L}{DPO} = -\mathbb{E}\left[\log \sigma\left(\beta \left(\log \frac{\pi\theta(y_w|x)}{\pi_{ref}(y_w|x)} - \log \frac{\pi_\theta(y_l|x)}{\pi_{ref}(y_l|x)}\right)\right)\right]$$

其中:

  • $y_w$:chosen(优选)回答
  • $y_l$:rejected(拒绝)回答
  • $\pi_\theta$:当前策略
  • $\pi_{ref}$:参考策略(通常是 SFT 模型)
  • $\beta$:KL 散度约束系数

PPO Loss

背景与动机

近端策略优化(Proximal Policy Optimization, PPO)通过裁剪重要性采样比率来限制策略更新幅度,防止策略崩溃。是 RLHF 训练的核心算法。

核心公式

$$\mathcal{L}_{PPO} = \mathbb{E}\left[\min\left(r_t(\theta) \hat{A}_t, \text{clip}(r_t, 1-\epsilon, 1+\epsilon)\hat{A}_t\right)\right]$$

其中:

  • $r_t = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{old}}(a_t|s_t)}$(重要性采样比率)
  • $\hat{A}_t$:优势函数估计
  • $\epsilon$:裁剪参数(通常 0.2)

裁剪机制图解

flowchart LR
    Advantage{"优势函数 A"}
    Advantage --> Positive["当 A > 0 (好动作)"]
    Advantage --> Negative["当 A < 0 (坏动作)"]
    Positive --> PositiveNormal["r 在 [1-ε, 1+ε] 内正常更新"]
    Positive --> PositiveClip["r > 1+ε 时停止奖励<br/>防止过度优化"]
    Negative --> NegativeNormal["r 在 [1-ε, 1+ε] 内正常惩罚"]
    Negative --> NegativeClip["r < 1-ε 时停止惩罚<br/>防止过度惩罚"]

GRPO Loss

背景与动机

分组相对策略优化(Group Relative Policy Optimization, GRPO)是 DeepSeek-Math 提出的算法,被 DeepSeek-R1 用于强化学习训练。它通过组内相对优势来消除对价值网络(Critic)的依赖。

核心思想

  • 对同一问题生成多个回答(组)
  • 在组内计算相对优势(而非绝对优势)
  • 无需训练 Critic 网络,简化训练流程

核心公式

1. 组内相对优势(Group-Relative Advantage)

对同一问题生成 $G$ 个回答,计算每个回答的奖励 $r_1, r_2, ..., r_G$,然后计算组内标准化优势:

$$\hat{A}_i = \frac{r_i - \text{mean}(\mathbf{r})}{\text{std}(\mathbf{r})}$$

2. GRPO 损失函数

$$\mathcal{L}{GRPO} = -\mathbb{E}\left[\frac{1}{G}\sum{i=1}^{G} \min\left(\rho_i \hat{A}_i, \text{clip}(\rho_i, 1-\epsilon, 1+\epsilon)\hat{A}i\right) - \beta \cdot \mathbb{D}{KL}\right]$$

其中:

  • $\rho_i = \frac{\pi_\theta(o_i|q)}{\pi_{\theta_{old}}(o_i|q)}$(重要性采样比率)
  • $\hat{A}_i$:组内相对优势
  • $\beta$:KL 散度惩罚系数
  • $\mathbb{D}_{KL}$:策略与参考策略的 KL 散度

与 PPO 的区别

特性PPOGRPO
优势估计需要 Critic 网络组内相对优势
额外网络需要 Value Head不需要
内存占用较高较低
适用场景通用 RL多候选生成场景

参数高效微调

LoRA

背景与动机

低秩适应(Low-Rank Adaptation, LoRA)通过在预训练权重旁添加低秩分解矩阵来实现参数高效微调。它冻结预训练权重,只训练少量参数,大大降低了微调成本。

核心思想:权重更新 $\Delta W$ 可以被低秩分解为 $B \cdot A$

核心公式

$$h = W_0 x + \Delta W x = W_0 x + BAx$$

其中:

  • $W_0 \in \mathbb{R}^{d \times k}$:冻结的预训练权重
  • $A \in \mathbb{R}^{r \times k}$:可训练,使用随机初始化
  • $B \in \mathbb{R}^{d \times r}$:可训练,初始化为零
  • $r \ll \min(d, k)$:低秩维度

关键设计:$B$ 初始化为零,使得初始状态 $BA = 0$,保证微调开始时模型行为不变。

张量形状流程图

flowchart TD
    X["输入 x: [batch, seq_len, in_features]"]
    X --> W0["W_0<br/>(frozen)"]
    X --> LoRA["LoRA分支"]
    W0 --> Base["W_0@x"]
    LoRA --> Dropout["Dropout"]
    Dropout --> A["lora_A: [rank, in_features]"]
    A --> B["lora_B: [out_features, rank]"]
    B --> Scaling["* scaling (α/r)"]
    Scaling --> Delta["B@A@x"]
    Base --> Add["+"]
    Delta --> Add
    Add --> Output["输出: [batch, seq_len, out_features]"]

参考文献

注意力机制

位置编码

归一化

前馈网络

训练方法

参数高效微调


Contributors

ckd0817

25 commits

chxu2000

1 commits

Languages

Python

60.0%

Jupyter Notebook

40.0%