Back

注意力机制深度技术解析:从数学原理到Transformer架构演进

注意力机制深度技术解析:从数学原理到 Transformer 架构演进

本文是注意力机制的完整技术指南,从数学推导到架构演进,适合有一定深度学习基础的读者。


一、为什么需要注意力机制?

1.1 序列建模的困境

在处理序列数据(文本、语音、时间序列)时,传统方法面临一个核心问题:如何捕捉长距离依赖?

RNN 的方案

h_t = f(h_{t-1}, x_t)

信息必须沿着时间步逐步传递,就像传话游戏:

x_1 → h_1 → h_2 → ... → h_{t-1} → h_t

问题

  • 距离越远,信息衰减越严重(梯度消失/爆炸)
  • 必须顺序计算,无法并行(慢)
  • 所有历史信息被压缩到一个固定大小的向量(瓶颈)

注意力的方案

直接计算任意两个位置的相关性,无需逐步传递

就像微信群里,你可以直接@任何人,不需要通过中间人传话。


二、文本如何进入模型:Tokenization

在理解注意力机制之前,我们需要先回答一个问题:文本是如何变成模型能处理的数字的?

2.1 为什么需要 Tokenization?

模型只能处理数字,不能直接处理文本。Tokenization 就是把文本切分成最小单元(Token),并映射为数字 ID 的过程。

原始文本:"我喜欢吃苹果"
    ↓ Tokenization
Token 序列:["我", "喜欢", "吃", "苹果"]
    ↓ 查表
ID 序列:[1234, 5678, 9012, 3456]
    ↓ 嵌入层
向量序列:[[0.1, 0.2, ...], [0.3, 0.4, ...], ...]

2.2 中英文分词的差异

语言 特点 分词方式
英文 词之间有空格 相对简单,按空格切分
中文 词之间无空格 需要专门的分词算法

英文示例

"I like eating apples"
→ ["I", "like", "eating", "apples"]  (按空格切分)

中文示例

"我喜欢吃苹果"
→ ["我", "喜欢", "吃", "苹果"]  (需要分词器识别词边界)

2.3 主流 Tokenization 算法

方法一:Word-level(词级别)

  • 每个词是一个 Token
  • 问题:词表太大(英文 10 万+,中文 5 万+),未登录词无法处理

方法二:Character-level(字符级别)

  • 每个字/字母是一个 Token
  • 问题:序列太长("hello" 要 5 个 Token),效率低

方法三:Subword-level(子词级别)— 当前主流

核心思想:常用词保持完整,罕见词拆分成子词。

BPE(Byte-Pair Encoding)算法

步骤 1:从字符级别开始
        "hello" → ["h", "e", "l", "l", "o"]

步骤 2:统计相邻字符共现频率
        "ll" 出现频率最高

步骤 3:合并最高频的相邻字符
        ["h", "e", "ll", "o"]

步骤 4:重复合并,直到达到目标词表大小
        ["he", "ll", "o"]  或  ["hello"](如果够频繁)

实际效果

文本 Token 序列 说明
"hello" ["hello"] 常用词,完整保留
"unhappiness" ["un", "happi", "ness"] 罕见词,拆分子词
"我喜欢吃苹果" ["我", "喜欢", "吃", "苹果"] 中文按词切分
"我喜欢吃苹果手机" ["我", "喜欢", "吃", "苹果", "手机"] 新词也能处理

2.4 Token ID 序列化

每个 Token 对应词表中的一个 ID:

# 词表示例(简化)
vocab = {
    "我": 1234,
    "喜欢": 5678,
    "吃": 9012,
    "苹果": 3456,
    "手机": 7890,
    "[PAD]": 0,    # 填充符
    "[UNK]": 1,    # 未知词
    "[BOS]": 2,    # 开始符
    "[EOS]": 3,    # 结束符
}

# 序列化
text = "我喜欢吃苹果"
tokens = ["我", "喜欢", "吃", "苹果"]
ids = [1234, 5678, 9012, 3456]

# 添加特殊符号
ids = [2, 1234, 5678, 9012, 3456, 3]  # [BOS] ... [EOS]

2.5 填充(Padding)与注意力掩码

问题:批次内序列长度不同,GPU 需要统一长度。

解决方案:短序列用 [PAD] 填充到最长序列的长度。

批次示例:
序列 1:"我喜欢吃苹果" → [2, 1234, 5678, 9012, 3456, 3, 0, 0]  (填充 2 个 PAD)
序列 2:"你好"       → [2, 1111, 2222, 3, 0, 0, 0, 0]          (填充 4 个 PAD)

注意力掩码:告诉模型哪些是真实 Token,哪些是填充。

掩码:[1, 1, 1, 1, 1, 1, 0, 0]  (1=真实,0=填充)

计算注意力时,填充位置的权重设为 -∞,softmax 后变为 0

三、注意力机制的数学原理

2.1 通用注意力框架

注意力机制的本质是一个加权求和过程:

Output=iαivi\text{Output} = \sum_{i} \alpha_i \cdot v_i

其中:

  • viv_i:值向量(Value),第 ii 个位置的实际内容
  • αi\alpha_i:注意力权重,表示"有多关注第 ii 个位置"
  • αi=1\sum \alpha_i = 1(权重归一化)

关键问题:如何计算 αi\alpha_i

2.2 三种注意力计算方式

方式一:加性注意力(Additive Attention)

αi=softmax(wTtanh(Wqq+Wkki))\alpha_i = \text{softmax}(w^T \tanh(W_q q + W_k k_i))

  • 通过一个单层神经网络计算相关性
  • 计算复杂度:O(nd2)O(n \cdot d^2)
  • 代表:Bahdanau Attention(2015)

方式二:点积注意力(Dot-Product Attention)

αi=softmax(qki)\alpha_i = \text{softmax}(q \cdot k_i)

  • 直接计算向量的内积
  • 计算复杂度:O(nd)O(n \cdot d)
  • 更简单、更快

方式三:缩放点积注意力(Scaled Dot-Product Attention)

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

  • 在点积基础上除以 dk\sqrt{d_k}(维度平方根)
  • 为什么需要缩放? 防止内积过大导致 softmax 梯度消失

推导

  • 假设 q,kq, k 的元素独立,均值为 0,方差为 1
  • qk=i=1dqikiq \cdot k = \sum_{i=1}^{d} q_i k_i 的方差为 dd
  • 除以 d\sqrt{d} 后,方差恢复为 1,数值稳定

2.3 Q、K、V 的直观理解

符号 名称 类比 作用
Q Query(查询) "我在找什么" 表示当前需要匹配的特征
K Key(键) "我有什么标签" 表示每个位置可被匹配的特征
V Value(值) "我的实际内容" 表示每个位置要传递的信息

图书馆类比

你要找一本关于"注意力机制"的书(Query)
    ↓
每本书都有标签(Key)
    ↓
你计算 Query 和每个 Key 的相似度(注意力权重)
    ↓
根据相似度,从所有书中提取内容(Value 加权求和)
    ↓
得到你需要的综合信息(Output)

三、自注意力(Self-Attention)

3.1 定义

自注意力 = 序列内部的注意力,Q、K、V 都来自同一个序列。

Self-Attention(X)=softmax(XWQ(XWK)Tdk)XWV\text{Self-Attention}(X) = \text{softmax}\left(\frac{XW_Q (XW_K)^T}{\sqrt{d_k}}\right) XW_V

其中:

  • XRn×dX \in \mathbb{R}^{n \times d}:输入序列(nn 个位置,每个 dd 维)
  • WQ,WK,WVW_Q, W_K, W_V:可学习的投影矩阵

3.2 计算步骤

输入:X = [x_1, x_2, ..., x_n]  (n个token,每个d维)

Step 1: 线性投影
Q = XW_Q  (n × d_k)
K = XW_K  (n × d_k)
V = XW_V  (n × d_v)

Step 2: 计算注意力分数
S = QK^T / √d_k  (n × n 的相似度矩阵)

Step 3: 归一化
A = softmax(S, dim=-1)  (每行和为1)

Step 4: 加权求和
Output = AV  (n × d_v)

3.3 计算示例

假设输入序列有 3 个 token:["The", "cat", "sat"],维度 dk=2d_k = 2

Q = [[0.1, 0.2],    K = [[0.3, 0.1],    V = [[0.5, 0.2],
     [0.3, 0.1],         [0.2, 0.4],         [0.1, 0.6],
     [0.2, 0.3]]         [0.1, 0.2]]         [0.3, 0.1]]

步骤 1:计算注意力分数 S=QKTS = QK^T

S = [[0.05, 0.10, 0.05],     ← "The" 对每个 token 的原始分数
     [0.11, 0.10, 0.08],     ← "cat" 对每个 token 的原始分数
     [0.09, 0.12, 0.08]]     ← "sat" 对每个 token 的原始分数

步骤 2:缩放 S/dk=S/2S/1.41S / \sqrt{d_k} = S / \sqrt{2} \approx S / 1.41

S_scaled = [[0.035, 0.071, 0.035],
            [0.078, 0.071, 0.057],
            [0.064, 0.085, 0.057]]

步骤 3:Softmax 归一化 得到注意力权重 A=softmax(Sscaled)A = \text{softmax}(S_{\text{scaled}})

A = [[0.33, 0.34, 0.33],     ← 每行和为 1
     [0.38, 0.32, 0.30],
     [0.34, 0.36, 0.30]]

步骤 4:加权求和 Output=AV\text{Output} = AV

Output = [[0.30, 0.30],    ← "The" 的表示(融合了所有 token 的信息)
          [0.28, 0.32],    ← "cat" 的表示
          [0.27, 0.31]]    ← "sat" 的表示

关键观察

  • 每个 token 的输出都融合了所有 token 的信息
  • 融合比例由注意力权重决定
  • 这是一个全连接的操作(每个位置看所有位置)

四、多头注意力(Multi-Head Attention)

4.1 动机

单个注意力头只能捕捉一种关系模式。但语言中有很多种关系:

"The cat sat on the mat because it was warm"

关系 1:指代关系(it → cat)
关系 2:空间关系(cat → on → mat)
关系 3:因果关系(sat → because → warm)

多头注意力 = 多组注意力并行,每组学习不同的关系模式。

4.2 公式

MultiHead(Q,K,V)=Concat(head1,...,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h) W_O

headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)

其中:

  • hh:头的数量(GPT-2 用 12 头,GPT-3 用 96 头)
  • WiQ,WiK,WiVW_i^Q, W_i^K, W_i^V:每个头的独立投影矩阵
  • WOW_O:输出投影矩阵

4.3 维度设计

模型维度:d_model = 768
头数:h = 12
每个头维度:d_k = d_v = d_model / h = 64

Q, K, V 维度:768 → 12 × 64
注意力矩阵:n × n(每个头独立计算)
输出拼接:12 × 64 → 768

为什么这样设计?

  • 保持总参数量不变(与单头相同维度)
  • 每个头专注于不同的子空间
  • 最后通过 WOW_O 融合所有头的信息

五、Transformer 架构详解

5.1 整体结构

Transformer
├── Encoder(编码器)× N 层
│   ├── Multi-Head Self-Attention
│   ├── Add & Norm(残差连接 + 层归一化)
│   ├── Feed-Forward Network
│   └── Add & Norm
└── Decoder(解码器)× N 层
    ├── Masked Multi-Head Self-Attention  ← 防止看到未来
    ├── Add & Norm
    ├── Multi-Head Cross-Attention        ← 看编码器输出
    ├── Add & Norm
    ├── Feed-Forward Network
    ── Add & Norm

5.2 关键组件

5.2.1 位置编码(Positional Encoding)

注意力机制本身没有位置信息(置换不变性),需要额外注入位置编码:

PE(pos,2i)=sin(pos/100002i/d)PE_{(pos, 2i)} = \sin(pos / 10000^{2i/d}) PE(pos,2i+1)=cos(pos/100002i/d)PE_{(pos, 2i+1)} = \cos(pos / 10000^{2i/d})

为什么用正弦函数?

  • 可以表示相对位置(PEpos+kPE_{pos+k} 可以表示为 PEposPE_{pos} 的线性函数)
  • 可以外推到更长的序列

现代替代方案

  • RoPE(旋转位置编码):LLaMA、Qwen 使用
  • ALiBi(线性偏置):不需要显式位置编码

5.2.2 层归一化(Layer Normalization)

LayerNorm(x) = γ · (x - μ) / √(σ² + ε) + β
  • 稳定训练,加速收敛
  • 位置:Pre-Norm(主流)vs Post-Norm(原始论文)

5.2.3 残差连接(Residual Connection)

Output = LayerNorm(x + Sublayer(x))
  • 解决深层网络梯度消失
  • 让信息可以直接流过

5.3 编码器 vs 解码器

特性 Encoder Decoder
自注意力 双向(看所有位置) 单向(掩码,只看过去)
交叉注意力 有(看编码器输出)
代表模型 BERT GPT、T5
应用场景 理解任务(分类、NER) 生成任务(翻译、摘要)

六、自回归:训练与预测流程

6.1 什么是自回归(Autoregressive)?

自回归 = 用过去预测未来,每次只生成一个 token,然后把新生成的 token 加回输入,继续生成下一个。

GPT 系列、LLaMA、Kimi 都是自回归模型
BERT 不是(它是双向的,能看到未来)

6.2 训练流程(Teacher Forcing)

目标

让模型学会:给定前面的词,预测下一个词

输入输出构造

假设训练文本是:

"我 喜欢 吃 苹果"

构造训练样本时,把同一个序列既当输入又当标签,但错开一位

输入(X) 标签(Y)
喜欢
我 喜欢
我 喜欢 吃 苹果
我 喜欢 吃 苹果 [EOS]

训练步骤

步骤 1:输入 "我"
        ↓
     模型预测下一个词的概率分布
        ↓
     标签是 "喜欢"
        ↓
     计算损失(交叉熵)

步骤 2:输入 "我 喜欢"
        ↓
     模型预测下一个词的概率分布
        ↓
     标签是 "吃"
        ↓
     计算损失

...(依次滑动窗口)

步骤 N:所有位置的损失求和
        ↓
     反向传播,更新模型参数

关键技巧:因果掩码(Causal Mask)

在 Transformer 中,用上三角掩码确保每个位置只能看到过去的 token:

注意力矩阵(4x4):
    我   喜欢  吃   苹果
我   ✓    ✗    ✗    ✗
喜欢  ✓    ✓    ✗    ✗
吃   ✓    ✓    ✓    ✗
苹果  ✓    ✓    ✓    ✓

✓ = 可以看到
✗ = 被掩码(看不到)

损失函数

L=t=1TlogP(xtx<t)\mathcal{L} = -\sum_{t=1}^{T} \log P(x_t | x_{<t})

即:对所有位置的下一个词预测计算交叉熵损失,然后求和。

6.3 预测/推理流程(自回归生成)

目标

给定一个提示词(Prompt),逐个生成后续 token,直到遇到结束符。

生成步骤

假设 Prompt 是 "今天天气",要生成后续内容:

步骤 1:输入 "今天天气"
        ↓
     模型输出下一个词的概率分布
        ↓
     采样/贪心选择 → "真"
        ↓
     输出:"今天天气真"

步骤 2:输入 "今天天气真"
        ↓
     模型输出下一个词的概率分布
        ↓
     采样/贪心选择 → "好"
        ↓
     输出:"今天天气真好"

步骤 3:输入 "今天天气真好"
        ↓
     模型输出下一个词的概率分布
        ↓
     采样/贪心选择 → ","
        ↓
     输出:"今天天气真好,"

...(循环,直到生成 [EOS] 或达到最大长度)

伪代码

def generate(model, prompt, max_length=100):
    tokens = tokenize(prompt)
    
    for _ in range(max_length):
        # 1. 模型前向传播
        logits = model(tokens)  # 形状:[seq_len, vocab_size]
        
        # 2. 取最后一个位置的输出
        next_logits = logits[-1]  # 形状:[vocab_size]
        
        # 3. 采样或贪心选择
        next_token = sample(next_logits)  # 或 argmax
        
        # 4. 检查是否结束
        if next_token == EOS:
            break
        
        # 5. 追加到输入
        tokens.append(next_token)
    
    return detokenize(tokens)

6.4 训练 vs 预测的核心区别

维度 训练 预测
输入 完整序列(一次喂入) 逐步增长(每次加一个 token)
输出 所有位置的预测(并行) 只取最后一个位置(串行)
标签 已知(Ground Truth) 未知(模型自己生成)
速度 快(GPU 并行) 慢(自回归循环)
暴露偏差 无(Teacher Forcing) 有(误差会累积)

6.5 推理优化:KV Cache

问题

自回归生成时,每生成一个新 token,都要重新计算前面所有 token 的注意力,大量重复计算

生成第 1 个 token:计算 1 个 token 的注意力
生成第 2 个 token:计算 2 个 token 的注意力(第 1 个重复计算)
生成第 3 个 token:计算 3 个 token 的注意力(前 2 个重复计算)
...

解决方案:KV Cache

缓存之前计算过的 K 和 V,每次只计算新 token 的 K 和 V。

步骤 1:输入 "今天天气"
        ↓
     计算 Q, K, V → 存入 KV Cache
        ↓
     输出 "真"

步骤 2:输入 "真"(只计算新 token)
        ↓
     计算 Q_new, K_new, V_new
        ↓
     拼接 KV Cache:K = [K_old, K_new]
        ↓
     输出 "好"

步骤 3:输入 "好"
        ↓
     计算 Q_new, K_new, V_new
        ↓
     拼接 KV Cache:K = [K_old, K_new]
        ↓
     输出 ","

效果

  • 计算量从 O(n2)O(n^2) 降到 O(n)O(n)
  • 生成速度提升 5-10 倍
  • 代价:显存占用增加(KV Cache 随序列线性增长)

6.6 采样策略

在预测时,如何从概率分布中选择下一个 token?

策略 方法 特点
贪心 argmax(P)\text{argmax}(P) 每次选概率最大的,确定性高,但可能重复
温度采样 P=softmax(P/T)P' = \text{softmax}(P/T) T<1 更确定,T>1 更随机
Top-K 从概率最大的 K 个中采样 避免低概率词
Top-P(核采样) 从累积概率 P 的词中采样 动态调整候选集
束搜索 保留 Beam 条最优路径 质量高,但速度慢

七、文本如何输出:解码

7.1 从 Token ID 到文本

模型输出的是 Token ID 序列,需要反向查表还原成文本。

# 解码示例
ids = [1234, 5678, 9012, 3456]
tokens = ["我", "喜欢", "吃", "苹果"]
text = "我喜欢吃苹果"

7.2 中英文解码差异

语言 解码方式 示例
英文 直接拼接,注意空格 ["I", "like", "apples"]"I like apples"
中文 直接拼接,无空格 ["我", "喜欢", "苹果"]"我喜欢苹果"

7.3 子词解码的特殊处理

使用 BPE 等子词分词时,需要处理词内连接符

# 英文 BPE 解码
tokens = ["un", "happi", "ness"]
# 去除连接符,合并
text = "unhappiness"

# 中文 BPE 解码
tokens = ["苹", "果"]
# 直接拼接
text = "苹果"

7.4 特殊 Token 的处理

Token 含义 解码时处理
[PAD] 填充符 忽略
[UNK] 未知词 显示为 [UNK] 或原样保留
[BOS] 开始符 忽略
[EOS] 结束符 停止解码

八、注意力的复杂度分析

8.1 标准注意力的瓶颈

操作 复杂度 说明
QKTQK^T O(n2d)O(n^2 \cdot d) 计算所有位置对的相似度
softmax O(n2)O(n^2) 归一化
AVAV O(n2d)O(n^2 \cdot d) 加权求和
总复杂度 O(n2d)O(n^2 \cdot d) 序列长度的平方

问题:当序列很长时(如 100K tokens),计算量和内存都不可接受。

8.2 长序列注意力优化方案

方案 核心思想 复杂度 代表模型
稀疏注意力 只看局部窗口 + 少量全局位置 O(nn)O(n \sqrt{n}) Longformer、BigBird
线性注意力 用核函数近似 softmax O(nd2)O(n \cdot d^2) Linformer、Performer
分块注意力 序列分块,块内全注意力 O(nm)O(n \cdot m) Reformer
状态空间模型 用递归代替注意力 O(n)O(n) Mamba、S4
KDA(因果线性注意力) 细粒度门控的线性注意力 O(n)O(n) Kimi K3

九、注意力的可视化与可解释性

9.1 注意力图(Attention Map)

注意力权重矩阵 ARn×nA \in \mathbb{R}^{n \times n} 可以可视化为热力图:

      The  cat  sat  on  the  mat
The [ 0.3  0.2  0.1  0.1  0.2  0.1 ]
cat [ 0.1  0.4  0.2  0.1  0.1  0.1 ]
sat [ 0.1  0.2  0.3  0.2  0.1  0.1 ]
...

解读

  • 对角线通常较高(token 关注自己)
  • 语法相关的 token 之间有较高权重
  • 指代关系可以通过注意力追踪

9.2 注意力头 specialization

不同注意力头会学习不同的模式:

头类型 特征 示例
位置头 关注相邻位置 捕捉局部语法
句法头 关注主谓/动宾关系 "cat" → "sat"
指代头 关注指代关系 "it" → "cat"
全局头 均匀关注所有位置 捕捉全局语境

十、注意力的前沿发展

10.1 线性注意力(Linear Attention)

核心思想:用核函数 ϕ\phi 近似 softmax,使计算可结合:

LinearAttention(Q,K,V)=ϕ(Q)(ϕ(K)TV)\text{LinearAttention}(Q, K, V) = \phi(Q) (\phi(K)^T V)

关键性质

  • 先计算 KTVK^T Vd×dd \times d 矩阵),再与 QQ 相乘
  • 复杂度从 O(n2d)O(n^2 d) 降到 O(nd2)O(n d^2)
  • dnd \ll n 时,近似线性

代表工作

  • Linformer(2020):低秩近似
  • Performer(2021):随机特征映射
  • KDA(2025):Kimi Delta Attention,细粒度门控(详见 8.5)

10.2 稀疏注意力(Sparse Attention)

核心思想:只计算部分位置的注意力,而非全部。

常见模式

1. 局部窗口:每个位置只看前后 w 个位置
2. 全局位置:特殊 token(如 [CLS])看所有位置
3. 随机连接:随机选择部分位置连接

代表模型:Longformer、BigBird、FlashAttention

10.3 FlashAttention

核心思想:不改变注意力数学,而是优化 GPU 内存访问

标准注意力:
QK^T → 写入 HBM → 读取 → softmax → 写入 HBM → 读取 → AV

FlashAttention:
分块计算,所有中间结果保持在 SRAM(片上内存)

效果

  • 速度提升 2-4 倍
  • 内存减少 5-20 倍
  • 数学结果完全相同(精确注意力)

10.4 因果注意力(Causal Attention)

核心思想:解码器只能看到过去的 token,不能看到未来。

实现方式:上三角掩码矩阵

Mij={0if ijif i<jM_{ij} = \begin{cases} 0 & \text{if } i \geq j \\ -\infty & \text{if } i < j \end{cases}

CausalAttention(Q,K,V)=softmax(QKT+Mdk)V\text{CausalAttention}(Q, K, V) = \text{softmax}\left(\frac{QK^T + M}{\sqrt{d_k}}\right)V

应用:所有自回归模型(GPT、LLaMA、Kimi)

10.5 KDA:Kimi Delta Attention(门控状态矩阵替代 softmax)

核心思想:KDA 不是简单的"线性注意力",而是用门控状态矩阵递推完全替代 softmax 注意力计算。它在保持线性复杂度的同时,通过通道级门控和 DPLR 并行结构弥补了传统线性注意力的表达能力不足。

KDA ≠ 标准线性注意力

  • 标准线性注意力:用核函数替代 softmax,无状态,表达能力弱
  • KDA:用门控状态矩阵替代 softmax,有状态 + 可并行,接近 softmax 的表达能力

状态更新方程

St=Diag(αt)St1+βtktvtS_t = \text{Diag}(\alpha_t) S_{t-1} + \beta_t k_t v_t^\top

组件 作用
StS_t 状态矩阵(替代 KV 缓存,固定大小)
Diag(αt)\text{Diag}(\alpha_t) 通道级遗忘门(每个特征维度独立遗忘率)
βt\beta_t 学习率(控制新信息写入速度)
ktvtk_t v_t^\top 外积(写入新记忆)

KDA vs 标准线性注意力

维度 标准线性注意力 KDA
核心方法 核函数替代 softmax 门控状态矩阵递推
是否有状态 无(每次重新计算) 有(StS_t 递推更新)
遗忘机制 通道级遗忘门 Diag(αt)\text{Diag}(\alpha_t)
并行训练 困难 DPLR 结构支持并行
表达能力 弱(无法精确检索) 强(接近 softmax)
复杂度 O(n)O(n) O(n)O(n)

三大创新

1. 通道级遗忘门(细粒度门控)

方案 遗忘粒度 效果
GDN(之前) 整个头共享一个 α\alpha 粗粒度,所有特征统一遗忘
KDA 每个特征维度独立 αi\alpha_i 细粒度,短文本快速遗忘,长文本缓慢遗忘

2. DPLR 并行计算结构

DPLR = Diagonal-Plus-Low-Rank(对角 + 低秩)

将状态转移矩阵分解为: 转移矩阵=对角矩阵+低秩矩阵\text{转移矩阵} = \text{对角矩阵} + \text{低秩矩阵}

  • 计算量从 4 次二级分块矩阵计算 → 2 次
  • 算子效率提升约 100%
  • 充分利用 GPU 张量核心

3. 混合架构(3:1 KDA-to-MLA)

层 1: KDA(门控状态矩阵)
层 2: KDA(门控状态矩阵)
层 3: KDA(门控状态矩阵)
层 4: MLA(全注意力)← 每 4 层插入 1 层全局注意力
层 5: KDA
...

设计哲学

  • 75% KDA:负责局部观察、模型压缩、加速(线性复杂度)
  • 25% MLA:负责全局建模、精准检索(全注意力)

性能对比

指标 标准注意力 标准线性注意力 KDA
复杂度 O(n2)O(n^2) O(n)O(n) O(n)O(n)
KV 缓存 随序列增长 固定状态矩阵
显存占用 100% 0% ~25%(状态矩阵)
表达能力 强(接近标准注意力)
百万 token 解码 1x 10x 6.3x

关键洞察

  • KDA 不是简单的线性注意力,而是门控状态矩阵方案
  • 通过通道级遗忘门弥补线性注意力的表达能力不足
  • 通过 DPLR 结构实现并行训练
  • 通过 3:1 混合架构兼顾效率和质量

类比理解

概念 标准注意力 KDA
记忆方式 记住所有历史(像录像) 压缩成状态矩阵(像摘要)
回忆方式 精确查找(softmax 权重) 状态递推(门控更新)
适用场景 短序列、需要精确检索 长序列、需要高效处理

10.6 MLA:Multi-head Latent Attention(多头潜在注意力)

核心思想:MLA 不是改变注意力计算,而是压缩 KV 缓存,把原始的 K 和 V 压缩成低维"潜在向量"。

MHA vs MLA 对比

概念 全称 核心特点
MHA Multi-Head Attention(标准多头注意力) 每个头独立存 K 和 V,显存占用大
MLA Multi-head Latent Attention(多头潜在注意力) 把 K 和 V 压缩成"潜在向量",显存占用极小

MHA 的问题

标准多头注意力(GPT-2/3、LLaMA 都在用):

头 1: Q1, K1, V1 → 存 KV 缓存
头 2: Q2, K2, V2 → 存 KV 缓存
...
头 h: Qh, Kh, Vh → 存 KV 缓存

问题:KV 缓存随序列长度线性增长,百万 token 时占用几十 GB 显存。

MLA 的压缩方案

压缩过程(推理时)

原始 K (维度 d_k) 
    ↓ 压缩矩阵 C_k
潜在向量 c_k^KV (维度 d_c,远小于 d_k)
    ↓ 存入 KV 缓存(极小)

还原过程(计算注意力时)

潜在向量 c_k^KV
    ↓ 解压矩阵
还原出近似的 K 和 V
    ↓ 计算注意力

数学表达K=WUKcKV+WSKxK = W_{UK} \cdot c^{KV} + W_{SK} \cdot x V=WUVcKV+WSVxV = W_{UV} \cdot c^{KV} + W_{SV} \cdot x

符号 含义
cKVc^{KV} 压缩后的潜在向量(存入缓存)
WUK,WUVW_{UK}, W_{UV} 解压矩阵(从缓存还原 K/V)
WSK,WSVW_{SK}, W_{SV} 跳过连接的补偿矩阵(减少压缩损失)

核心对比

维度 MHA(标准) MLA(潜在)
KV 缓存内容 原始 K 和 V 压缩后的潜在向量 cKVc^{KV}
缓存维度 dk+dvd_k + d_v(大) dcd_c(小,通常只有原来的 1/10)
显存占用 降低 90%+
计算精度 精确 有损压缩(但通过跳过连接补偿)
代表模型 GPT-3/4、LLaMA、Qwen DeepSeek-V2/V3、Kimi K3

类比理解

概念 类比
MHA 把每页书都原样复印存起来(占地方,但查阅精确)
MLA 把每页书压缩成摘要存起来(省地方,查阅时再展开)
KDA 不存书,只记一个"读书笔记"(最省地方,但可能丢细节)

10.7 现代架构演进总结

2017: Transformer(MHA)
    ↓ 显存瓶颈
2023: DeepSeek-V2(MLA)— 压缩 KV 缓存
    ↓ 长序列瓶颈
2025: Kimi Linear(KDA + MLA 混合)— 门控状态矩阵 + 压缩全注意力
    ↓ 未来
2026+: 一步生成 + 物理一致性 + 世界模型

Kimi K3 的终极方案

  • 75% KDA(线性注意力)— 完全不要 KV 缓存,O(n)O(n) 复杂度
  • 25% MLA(压缩全注意力)— 保留全局精确注意力,显存降低 90%

兼顾效率和质量的平衡之道


十一、注意力机制的工程优化

11.1 KV Cache(推理优化)

问题:自回归生成时,每一步都要重新计算所有历史 token 的 K、V。

解决:缓存历史 token 的 K、V,只计算新 token。

Step 1: 计算 x_1 的 Q, K, V → 缓存 K_1, V_1
Step 2: 计算 x_2 的 Q, K, V → 缓存 K_2, V_2,复用 K_1, V_1
Step 3: 计算 x_3 的 Q, K, V → 缓存 K_3, V_3,复用 K_1, V_2, K_2, V_2
...

效果:推理速度提升 10-50 倍

11.2 分组查询注意力(GQA)

问题:多头注意力中,每个头独立的 K、V 占用大量内存。

解决:多个 Q 头共享一组 K、V。

MHA(多头注意力):32 个 Q 头,32 个 K 头,32 个 V 头
GQA(分组查询):  32 个 Q 头,8 个 K 头,8 个 V 头(每 4 个 Q 共享 1 组 KV)
MQA(多查询):    32 个 Q 头,1 个 K 头,1 个 V 头(所有 Q 共享 1 组 KV)

效果:KV Cache 减少 4-32 倍,推理速度大幅提升

代表模型:LLaMA 2/3(GQA)、Falcon(MQA)


十二、总结与展望

12.1 注意力机制的核心价值

维度 贡献
性能 解决长距离依赖,超越 RNN
效率 支持并行计算,训练速度提升
可解释性 注意力权重可可视化
通用性 适用于 NLP、CV、音频、多模态

12.2 未来方向

  1. 线性注意力:突破 O(n2)O(n^2) 复杂度瓶颈(Kimi K3 的 KDA)
  2. 动态稀疏:根据输入动态选择关注位置
  3. 注意力+SSM 混合:结合注意力和状态空间模型的优势
  4. 高效推理:KV Cache 压缩、量化、蒸馏

12.3 一句话总结

注意力机制让模型学会了"看重点",Transformer 把这个能力发挥到极致,大模型时代由此开启。


附录:关键公式汇总

A.1 缩放点积注意力

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

A.2 多头注意力

MultiHead(Q,K,V)=Concat(head1,...,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h) W_O

A.3 位置编码(正弦)

PE(pos,2i)=sin(pos/100002i/d)PE_{(pos, 2i)} = \sin(pos / 10000^{2i/d})

A.4 因果掩码

Mij={0if ijif i<jM_{ij} = \begin{cases} 0 & \text{if } i \geq j \\ -\infty & \text{if } i < j \end{cases}

A.5 线性注意力(核函数近似)

LinearAttention(Q,K,V)=ϕ(Q)(ϕ(K)TV)\text{LinearAttention}(Q, K, V) = \phi(Q) (\phi(K)^T V)


参考论文

  1. Vaswani et al. "Attention Is All You Need" (2017)
  2. Bahdanau et al. "Neural Machine Translation by Jointly Learning to Align and Translate" (2015)
  3. Dosovitskiy et al. "An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale" (2020)
  4. Dao et al. "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness" (2022)
  5. Kimi Team. "Kimi Linear: Hybrid Linear Attention with 3:1 KDA-to-MLA Ratio" (2025)

评论