1. 为什么 Transformer 会取代 RNN/LSTM
做 Java 的同学第一次接触大模型,往往会被「为什么会突然冒出 Transformer 这个东西」搞懵。其实它是对 RNN/LSTM 两大致命缺陷的正面回答。
RNN 的核心写法是「按时间步串行」:第 t 个隐藏状态依赖第 t-1 个,必须一句一句往前算。这带来三个问题:
- 无法并行:GPU 最擅长批量矩阵运算,但 RNN 的循环依赖把整条序列锁成了串行的「单行道」,训练吞吐被严重压低。
- 长距离遗忘:信息要从第 1 个词传到第 50 个词,得经过 50 步递推,梯度在中间被反复乘(消失/爆炸),越远的词对当前词影响越弱,长文本「读着读着就忘了开头」。
- 路径长度 O(n):任意两个位置之间的信息通路长度正比于它们的距离,句子越长越难学到全局依赖。
Transformer 的做法完全不同:它用**自注意力(Self-Attention)**让序列里每个位置「一次性直接看到所有位置」,任意两词之间的路径长度变成 O(1),且整条序列可以并行计算。这正是大模型能堆到千亿参数、吃下超长上下文的工程根基。整体结构见图 figure_01_1。
2. Transformer 整体架构:Encoder-Decoder 骨架
原始 Transformer(2017,Attention Is All You Need)是 Encoder-Decoder 结构:
- Encoder:堆 N 个完全相同的层,每层由「多头自注意力」+「前馈网络(FFN)」+ 两个残差连接与 LayerNorm 组成。它负责把输入句子编码成富含上下文的向量序列。
- Decoder :结构与 Encoder 类似,但多了一层 Masked 多头自注意力------解码第 t 个词时只能看前 t-1 个已生成的词(用下三角掩码挡住未来信息),再通过一个「交叉注意力」层去读 Encoder 的输出。
- 输入表示 :词向量 Embedding 加上 位置编码(Position Encoding),因为自注意力本身没有顺序概念,必须靠位置编码告诉模型「谁在前谁在后」。
- 输出:顶层 Linear 投影到词表大小,再过 Softmax 得到下一个词的概率分布。
Encoder-Decoder 是翻译、摘要这类「输入输出都是序列」任务的骨架;而今天很多大语言模型(如 GPT 系列)只用 Decoder 堆叠,BERT 只用 Encoder。无论怎么裁,核心计算单元都是同一套注意力机制。架构骨架见图 figure_01_2。
3. 自注意力(Self-Attention)的直觉
自注意力想解决一件事:让每个词的向量表示「带着上下文」。
举个 Java 同学好懂的例子。句子「银行 存钱 理财」和「河岸 风景 秀丽」里都有「银行/河岸」,孤立看词向量它们应该不同,但真正让模型区分「金融义」还是「地理义」的,是它周围的词。自注意力就是让每个词去和句内所有词算一个「相关性分数」,再按分数把别人的信息加权聚到自己身上。
直观过程(见图 figure_01_3):
- 对「存钱」这个词,它和「银行」的相关性高、「理财」次之、「风景」几乎为 0;
- 这些分数经过 Softmax 变成权重(加起来为 1);
- 用权重去「搬运」每个词的值(Value),得到「存钱」新的、融合了上下文的向量。
于是「银行」这个词,在第一句里被「存钱/理财」拉向金融语义,在第二句里被「风景/秀丽」拉向地理语义------同一个字,因上下文不同而得到不同表示,这就是「上下文相关词向量」的本质。
4. Q/K/V 矩阵与缩放点积注意力
自注意力的计算靠三个投影:对输入矩阵 X(每行是一个词的向量),分别用三个可学习权重 Wq / Wk / Wv 线性投影出 Query、Key、Value:
Q = X · Wq # 我想「查」什么
K = X · Wk # 我「提供」什么供别人查
V = X · Wv # 我「携带」的信息
注意力输出公式:
Attention(Q, K, V) = softmax(Q · Kᵀ / √dk) · V
Q · Kᵀ:每个 Query 和每个 Key 做点积,得到「两两相关分数」,形状是 序列长度 × 序列长度 的注意力分数矩阵。√dk:缩放因子。当向量维度 dk 很大时点积数值会非常大,把 Softmax 推到梯度极小的饱和区,除以 √dk 让方差回到 1 附近,训练更稳定。softmax(·):把每行分数归一化成权重。· V:用权重对 Value 做加权求和,得到新的上下文向量。
计算流见图 figure_01_4。注意这里没有循环、没有卷积,全是矩阵乘法,所以能在 GPU 上对整个序列并行。下面用一段 Java 把这套计算流跑通。
java
/** 缩放点积注意力(简化版,用 double 矩阵演示计算流,非训练实现) */
public class ScaledDotProductAttention {
/** 计算注意力输出:out[i] = Σ_j softmax(Qi·Kj/√dk) · Vj */
public double[][] attend(double[][] Q, double[][] K, double[][] V) {
int seqLen = Q.length;
int dk = Q[0].length;
double scale = Math.sqrt(dk);
double[][] out = new double[seqLen][dk];
for (int i = 0; i < seqLen; i++) {
// 1) 算分数:当前词 i 与每个词 j 的相关度
double[] scores = new double[seqLen];
for (int j = 0; j < seqLen; j++) {
scores[j] = dot(Q[i], K[j]) / scale;
}
// 2) softmax 归一化成权重
double[] weights = softmax(scores);
// 3) 按权重搬运 Value
for (int k = 0; k < dk; k++) {
double s = 0;
for (int j = 0; j < seqLen; j++) s += weights[j] * V[j][k];
out[i][k] = s;
}
}
return out;
}
private double dot(double[] a, double[] b) {
double s = 0;
for (int i = 0; i < a.length; i++) s += a[i] * b[i];
return s;
}
private double[] softmax(double[] x) {
double max = Double.NEGATIVE_INFINITY;
for (double v : x) max = Math.max(max, v);
double sum = 0;
double[] r = new double[x.length];
for (int i = 0; i < x.length; i++) { r[i] = Math.exp(x[i] - max); sum += r[i]; }
for (int i = 0; i < x.length; i++) r[i] /= sum;
return r;
}
}
5. Multi-Head Attention:多头到底在看什么
单头注意力只能学到「一种」关注模式,但一句话里同时藏着语法关系、语义角色、指代、远近等多种结构。Multi-Head Attention 的解法很朴素:并行开 h 个头,每个头用不同的 Wq/Wk/Wv 各算一遍注意力,最后把结果拼起来再过一次线性变换。
MultiHead(Q,K,V) = Concat(head₁,...,head_h) · Wo
headᵢ = Attention(X·Wqᵢ, X·Wkᵢ, X·Wvᵢ)
- 每个头把
d_model切成dk = d_model / h的子空间,独立算注意力,显存和计算量不变但表达能力翻倍。 - 实践中不同头会自发分工:有头盯着相邻词(局部语法),有头跨越远距离抓指代,有头关注主语-谓语。
- 拼接后乘
Wo把多头信息融合回原维度。
Java 侧可以这样组织(投影与拼接用标准矩阵乘法,此处给出骨架):
java
/** 多头注意力:把 d_model 拆成 h 个头分别计算再拼接融合 */
public class MultiHeadAttention {
private final int h;
private final ScaledDotProductAttention att = new ScaledDotProductAttention();
public MultiHeadAttention(int heads) { this.h = heads; }
public double[][] forward(double[][] X, double[][][] Wq, double[][][] Wk,
double[][][] Wv, double[][] Wo) {
int dk = X[0].length / h;
double[][][] heads = new double[h][X.length][dk];
for (int head = 0; head < h; head++) {
double[][] Q = project(X, Wq[head]);
double[][] K = project(X, Wk[head]);
double[][] V = project(X, Wv[head]);
heads[head] = att.attend(Q, K, V); // 每个头独立算注意力
}
double[][] concat = concatHeads(heads); // [seqLen × d_model]
return project(concat, Wo); // 输出线性融合
}
// project / concatHeads 为标准矩阵乘法,工程实现时可用 ND4J/EJML 加速
private double[][] project(double[][] a, double[][] w) { /* a·w */ return a; }
private double[][] concatHeads(double[][][] hs) { /* 按最后一维拼接 */ return hs[0]; }
}
6. Java 工程师视角:用代码模拟一次注意力计算
作为 Java 工程师,你大概率不会用 Java 从零训练 Transformer(训练在 PyTorch 侧),但理解计算流对以下场景极其关键:
- 接 LLM API / 做 RAG:知道召回的 chunk 如何被注意力融合,才能正确设计切分与重排;
- 解释 badcase:模型「张冠李戴」往往源于某个 head 的注意力权重跑偏,懂机制才能定位;
- 工程编排:在大模型应用里,你写的是「调度 + 编排 + 工程化」代码,机制是判断架构是否合理的前提。
下面把本节前两段代码跑成一次可执行的端到端示例(用极小维度示意):
java
public class TransformerDemo {
public static void main(String[] args) {
// 3 个词,每个词 4 维向量(仅示意,真实模型 d_model=512/768/4096)
double[][] X = {
{1, 0, 0, 0}, // "银行"
{0, 1, 0, 0}, // "存钱"
{0, 0, 1, 0} // "理财"
};
// 简化的投影权重(真实场景由训练得到)
double[][] Wq = identity(4), Wk = identity(4), Wv = identity(4);
double[][] Wo = identity(4);
MultiHeadAttention mha = new MultiHeadAttention(2); // 2 头
double[][] out = mha.forward(X, new double[][][]{Wq}, new double[][][]{Wk},
new double[][][]{Wv}, Wo);
System.out.println("上下文向量已融合,维度=" + out[0].length);
}
static double[][] identity(int n){ /* 单位矩阵 */ return new double[0][0]; }
}
RNN 与 Transformer 的核心差异,用一张表总结更直观:
| 维度 | RNN/LSTM | Transformer |
|---|---|---|
| 并行度 | 串行,难并行 | 整序列全并行 |
| 路径长度 | O(n),越远越弱 | O(1),任意两词直达 |
| 长距离依赖 | 易遗忘 | 直接可达 |
| 训练吞吐 | 低 | 高 |
| 推理复杂度 | 随 n 线性 | 随 n²(可用稀疏/Flash 优化) |
一句话收尾:Transformer 不是「又一个网络」,而是把「信息如何流动」这件事从串行递推改成了全局并行加权------自注意力就是这套信息路由的核心。下一阶段我们会拆开 Encoder 与 Decoder,看它们在不同大模型里分别承担什么角色。