一句话总结
Transformer 用 Self-Attention 机制让每个词同时看到句子中所有其他词,彻底解决了 RNN/LSTM 的顺序处理瓶颈.
理解 Attention 的直觉(查字典)比记住公式更重要--它是 BERT, GPT 等一切现代模型的核心.
前置回顾
阶段二我们用 LSTM 跑通了文本分类,也体会到了它的两个瓶颈:训练慢(必须逐词计算,无法并行)和长距离遗忘(句首信息传到句尾时已经衰减). 本篇讲 Transformer 如何一次性解决这两个问题.
LSTM 的瓶颈:逐词排队
LSTM 处理一句话的方式像流水线--必须先处理第 1 个词,才能处理第 2 个,以此类推:
|
两个致命问题:
- 速度:GPU 擅长并行计算,但 LSTM 强制串行. 一个 batch 里 100 条文本,每条 20 个词,需要 20 轮串行步骤
- 距离:"管理员"的信息要经过 5 步传递才能到达"用户". 每一步都有信息损耗,句子越长,开头的信息越容易被"冲淡"
LSTM = 传话游戏. 一排人依次传话,"管理员禁言了发广告的用户"传到最后一个人可能变成"有人被禁言了". 中间每个人都会丢失或扭曲一部分信息.
Transformer = 开会. 所有人坐在一个圆桌旁,每个人同时能听到所有其他人说的话,不存在传话损耗.
Self-Attention:每个词都能直接看到每个词
Transformer 的核心机制是 Self-Attention(自注意力). 核心思想极简:对于句子中的每一个词,计算它和其他所有词的"相关度",然后按相关度加权融合信息.
用一个直觉模型来理解--查字典:
|
结果:"禁言"这个词的输出向量里已经融合了"谁执行""对谁执行""为什么执行"的信息--不需要传话,直接拿到.
Self-Attention = Go 的 map 查找. Q 是 key,K 是 map 中每个 entry 的 key,V 是 value. 点积就是计算两个 key 的匹配程度. 不同的是 Attention 做的是模糊匹配--不是精确命中一个 entry,而是按匹配度从所有 entry 中加权取值.
注意力权重的实际含义
把注意力权重画成热力图,可以直观看到模型在"关注"什么:
|
这就是 Attention 的"可解释性":可以看到模型在做决策时,每个词在关注哪些上下文.
多头注意力:从不同角度看
一组 Q/K/V 只能捕获一种关系模式. 但一个词在不同维度上有不同的关系:
- "禁言"和"管理员"的关系:施动者关系
- "禁言"和"用户"的关系:受动者关系
- "禁言"和"发广告"的关系:原因关系
多头注意力(Multi-Head Attention) 就是同时运行多组独立的 Q/K/V,每组(称为一个"头")学习一种关注模式:
|
维度拆分
BERT-base 的隐藏维度是 768,12 个头意味着每个头操作 768/12 = 64 维. 这不是额外增加计算量,而是把同样的维度拆成 12 份并行处理. 总计算量和单头 768 维差不多,但能捕获更多样的关系模式.
位置编码:告诉模型词的顺序
Self-Attention 有一个问题:它是集合运算,不区分词的位置. "管理员禁言了用户"和"用户禁言了管理员"在纯 Attention 计算中是一样的--因为每个词都能看到所有其他词,和顺序无关.
解决方案是位置编码(Positional Encoding):给每个位置一个固定的向量,加到 Embedding 上:
|
原始 Transformer 用正弦/余弦函数生成位置编码(波长从短到长覆盖不同频率). BERT 换了一种更简单的方式:直接训练一个位置 Embedding 表(和词 Embedding 表一样,位置 ID → 向量,通过训练学出来).
最大长度限制
BERT 的位置 Embedding 表只有 512 行,所以最多处理 512 个 token. 超过就截断. 对群聊消息来说 512 绰绰有余(大部分消息不到 50 个字),但如果处理长文档就需要 Longformer 等变体.
为什么 Transformer 可以并行
这是 Transformer 相比 LSTM 最大的工程优势. 关键在于:Self-Attention 中所有词的 Q/K/V 计算是同时进行的.
|
| 特性 | LSTM | Transformer |
|---|---|---|
| 计算方式 | 逐步串行 | 全部并行 |
| 长距离依赖 | 信息衰减,需要门控补救 | 直接 Attention,无衰减 |
| GPU 利用率 | 低(串行瓶颈) | 高(全是矩阵运算) |
| 训练速度(同等数据) | 慢 3-10x | 快 |
| 可解释性 | 隐藏状态难以解释 | 注意力权重可视化 |
LSTM 训练 = 单核 CPU 跑 for 循环. 每次迭代依赖上一次结果,无法并行.
Transformer 训练 = GPU 跑矩阵乘法. 所有元素同时计算,硬件利用率拉满. 这就是为什么 Transformer 能在更大的数据集上训练更大的模型--算力效率有质的提升.
Encoder-Decoder 架构
原始 Transformer("Attention Is All You Need"论文)是为翻译任务设计的,包含两个部分:
|
- Encoder:读入完整的输入序列,用 Self-Attention 编码上下文信息. 每个位置能看到所有其他位置(双向)
- Decoder:逐步生成输出序列. 每一步只能看到已生成的部分(用 Mask 遮住未来位置),同时通过 Cross-Attention 查看 Encoder 的输出
- Cross-Attention:Decoder 用自己的 Q 去查 Encoder 的 K/V. 直觉是"生成每个目标词时,去源句子里找相关信息"
不同任务使用架构的不同部分:
| 模型 | 使用部分 | 适合任务 |
|---|---|---|
| BERT | 只用 Encoder | 分类,实体识别,阅读理解 |
| GPT | 只用 Decoder | 文本生成,对话 |
| T5, BART | Encoder + Decoder | 翻译,摘要,问答 |
为什么群聊管理用 BERT
我们的任务是分类(一句话 → 一个标签),不需要生成文本. BERT 只用 Encoder,对输入做双向理解,天然适合分类任务. GPT 是单向的(只能看到左边的词),对分类任务来说浪费了右边的上下文信息.
一层 Transformer 的完整流程
把前面的概念串起来,一层 Transformer Encoder 做了什么:
|
BERT-base 把这个结构叠了 12 层. 每一层的输出是下一层的输入. 12 层之后,每个位置的向量已经编码了整句话从局部搭配到全局语义的多层次信息.
残差连接的重要性
没有残差连接,12 层叠加后梯度会指数级衰减,模型训练不动. 残差连接让梯度可以"走捷径"直接传回浅层. 如果你熟悉 ResNet--同一个思路. 看到 x + f(x) 的结构就是残差.
从 LSTM 到 Transformer:群聊场景的实际影响
回到群聊对骂检测的场景,看看 Transformer 带来了什么实质改善:
|
| 场景 | LSTM 的困难 | Transformer 的优势 |
|---|---|---|
| 长句(30+ 字) | 开头信息衰减 | 直接 Attention,无距离限制 |
| 反讽("你可真行啊") | 语气词在句尾,和前文距离远 | 同时看到所有词,捕获反讽模式 |
| 引用("他说xxx") | 难以区分引用内容和说话者态度 | 可以学到不同的注意力模式区分引用和评论 |
| 批量推理 | 串行,每秒处理量有限 | 并行,同等硬件吞吐量高数倍 |
"Attention Is All You Need" 的关键洞察
2017 年 Google 发表的这篇论文,核心贡献不是发明了 Attention(Attention 机制在此之前已经存在),而是证明了:
- 不需要 RNN/CNN:纯 Attention 堆叠就能达到甚至超过 RNN 的效果
- 并行训练:去掉了序列依赖,训练速度大幅提升,使得在更大数据集上训练成为可能
- 统一架构:同一个架构(Encoder-Decoder)可以处理翻译,摘要,问答等多种任务
这篇论文之后,NLP 领域几乎全面转向 Transformer. BERT(2018)取了 Encoder,GPT(2018)取了 Decoder,后续几乎所有主流模型都是 Transformer 的变体.
快速回顾
- Self-Attention:每个词通过 Q/K/V 机制直接查看所有其他词,用相关度加权融合信息
- 多头注意力:12 个头并行,每个头捕获不同类型的关系(主谓/动宾/语气等)
- 位置编码:给每个位置加一个向量,补偿 Attention 不区分顺序的问题
- 并行优势:所有位置同时计算,GPU 利用率远高于 LSTM
- Encoder vs Decoder:BERT 用 Encoder(双向,适合分类),GPT 用 Decoder(单向,适合生成)
- 残差连接:
x + f(x)结构让深层网络可训练
动手练习
- 可视化注意力:用
bertviz库加载bert-base-chinese,输入一句群聊消息,可视化 12 层 × 12 头的注意力热力图,观察不同头关注的模式 - 对比推理速度:分别用 LSTM 和 BERT 对 1000 条测试文本做推理,记录耗时. 观察 batch_size 从 1 增加到 32 时两者的速度变化差异
- 位置编码实验:用
model.embeddings.position_embeddings.weight取出位置编码矩阵,计算相邻位置和远距离位置的余弦相似度,验证"近距离位置编码更相似"的直觉 - 截断影响测试:构造一些超长文本(> 200 字),分别截断到 64 / 128 / 256 / 512 token,观察分类效果是否有显著差异