Multi-Head Attention 从数学到代码:TypeScript 并行多头实现完全指南
一、问题引入:一个「头」够吗?
单头 Self-Attention 在算 softmax(Q·Kᵀ/√d_k)·V 时,对所有 token 对只产生一种注意力模式。但这远远不够。看这句话:
「the animal didn't cross the street because it was too tired」
这里「it」应该关注「animal」(语义关系)、「tired」(因果关系),同时还需要关注当前句子的语法结构(句法关系)。单一的注意力分布无法同时捕捉这三种不同类型的依赖。
| 注意力类型 | 单头表现 | 多头表现 |
|---|---|---|
| 语义关系 (it→animal) | 可能覆盖 | ✓ 独立头专门捕捉 |
| 句法结构 (it→主语位置) | 被语义权重稀释 | ✓ 另一头专门捕捉 |
| 长距离依赖 (it→开头的 the) | 权重分散 | ✓ 再一头捕捉 |
Multi-Head Attention 的解决方案很直接:并行运行多个独立的 Self-Attention,每个头关注不同的表示子空间,最后把所有头的结果 Concat 起来。
二、核心概念速览
多头注意力的计算流水线:
X [n, d_model]
↓
分头: 每个头有独立的 W_Q^i, W_K^i, W_V^i
↓
┌─────────────────────────────────────┐
│ Head_1 │ Head_2 │ ... │ Head_h │ 并行计算
│ Q₁K₁ᵀV₁│ Q₂K₂ᵀV₂│ │ Q_hK_hᵀV│
└─────────────────────────────────────┘
↓
Concat: [Head_1, Head_2, ..., Head_h] → [n, h·d_v]
↓
Linear: W_O [h·d_v, d_model] → [n, d_model]
| 参数 | 论文默认值 | 含义 |
|---|---|---|
| h (头数) | 8 | 并行注意力头数 |
| d_model | 512 | 模型总维度 |
| d_k = d_v | 64 = 512/8 | 每个头的 Q/K/V 维度 |
| W_O | [512, 512] | 多头拼接后的输出投影矩阵 |
关键公式:
MultiHead(Q, K, V) = Concat(head_1, ..., head_h) · W_O
head_i = Attention(Q·W_Q^i, K·W_K^i, V·W_V^i)
三、完整可运行代码
/**
* Multi-Head Attention — TypeScript 完整实现
* 参考: "Attention Is All You Need" Section 3.2.2
*
* h = 8 头,每个头独立 Q/K/V 投影,Concat 后经 W_O 输出
*/
// ---------- 矩阵工具(复用单头实现的基础函数)----------
function randMatrix(rows: number, cols: number): number[][] {
return Array.from({ length: rows }, () =>
Array.from({ length: cols }, () => Math.random() * 0.02 - 0.01)
);
}
function matMul(A: number[][], B: number[][]): number[][] {
const m = A.length, n = A[0].length, p = B[0].length;
if (n !== B.length) throw new Error(`维度不匹配`);
const result = Array.from({ length: m }, () => new Array(p).fill(0));
for (let i = 0; i < m; i++) {
for (let k = 0; k < n; k++) {
const aik = A[i][k];
if (aik === 0) continue;
const rowB = B[k];
for (let j = 0; j < p; j++) result[i][j] += aik * rowB[j];
}
}
return result;
}
function transpose(M: number[][]): number[][] {
const rows = M.length, cols = M[0].length;
const result = Array.from({ length: cols }, () => new Array(rows).fill(0));
for (let i = 0; i < rows; i++) {
for (let j = 0; j < cols; j++) result[j][i] = M[i][j];
}
return result;
}
function softmaxRows(matrix: number[][]): number[][] {
return matrix.map(row => {
const maxVal = Math.max(...row);
const exps = row.map(v => Math.exp(v - maxVal));
const sumExp = exps.reduce((a, b) => a + b, 0);
return exps.map(v => v / sumExp);
});
}
// ---------- 单头 Attention 函数(无状态版本)----------
function singleHeadAttention(
Q: number[][],
K: number[][],
V: number[][],
mask?: number[][]
): { output: number[][]; weights: number[][] } {
const d_k = Q[0].length;
const seqLen = Q.length;
const scale = Math.sqrt(d_k);
const scores = matMul(Q, transpose(K));
// 缩放 + Mask
const scaled = scores.map((row, i) =>
row.map((v, j) => {
const s = v / scale;
if (mask && mask[i] && mask[i][j] === 0) return -1e9;
return s;
})
);
const weights = softmaxRows(scaled);
const output = matMul(weights, V);
return { output, weights };
}
// ---------- Multi-Head Attention ----------
interface MHAConfig {
d_model: number; // 模型总维度
h: number; // 注意力头数
d_k: number; // 每个头的 K/Q 维度
d_v: number; // 每个头的 V 维度
}
class MultiHeadAttention {
// 每个头有独立的 3 个权重矩阵
private WQ: number[][][]; // [h][d_model, d_k]
private WK: number[][][]; // [h][d_model, d_k]
private WV: number[][][]; // [h][d_model, d_v]
private WO: number[][]; // [h * d_v, d_model] 输出投影
private config: MHAConfig;
constructor(config: MHAConfig) {
this.config = config;
const { h, d_model, d_k, d_v } = config;
this.WQ = Array.from({ length: h }, () => randMatrix(d_model, d_k));
this.WK = Array.from({ length: h }, () => randMatrix(d_model, d_k));
this.WV = Array.from({ length: h }, () => randMatrix(d_model, d_v));
this.WO = randMatrix(h * d_v, d_model);
}
/**
* 前向传播
* @param X [seq_len, d_model] 输入
* @param mask 可选的注意力掩码
*/
forward(
X: number[][],
mask?: number[][]
): {
output: number[][];
allHeadWeights: number[][][]; // [h][seq_len, seq_len]
} {
const { h } = this.config;
const seqLen = X.length;
// 并行计算所有头(实际实现中可向量化)
const headOutputs: number[][][] = [];
const allHeadWeights: number[][][] = [];
for (let i = 0; i < h; i++) {
// 每个头独立投影
const Q = matMul(X, this.WQ[i]);
const K = matMul(X, this.WK[i]);
const V = matMul(X, this.WV[i]);
const { output, weights } = singleHeadAttention(Q, K, V, mask);
headOutputs.push(output);
allHeadWeights.push(weights);
}
// Concat: [seqLen, d_v] * h → [seqLen, h * d_v]
const concat: number[][] = Array.from({ length: seqLen }, () => []);
for (let pos = 0; pos < seqLen; pos++) {
for (let i = 0; i < h; i++) {
concat[pos].push(...headOutputs[i][pos]);
}
}
// 最终线性投影
const output = matMul(concat, this.WO); // [seqLen, d_model]
return { output, allHeadWeights };
}
}
// ---------- 测试 ----------
function test() {
const config: MHAConfig = {
d_model: 512,
h: 8,
d_k: 64,
d_v: 64,
};
const mha = new MultiHeadAttention(config);
const seqLen = 10;
const X = randMatrix(seqLen, config.d_model);
// 不带 Mask(Encoder 模式)
const { output, allHeadWeights } = mha.forward(X);
console.log(`MHA 输入: [${seqLen}, ${config.d_model}]`);
console.log(`MHA 输出: [${output.length}, ${output[0].length}]`);
console.log(`注意力头数: ${allHeadWeights.length}`);
// 对比:不同头的注意力模式差异
console.log('\n各头 Token0→所有Token 注意力分布对比:');
for (let i = 0; i < Math.min(3, config.h); i++) {
const w = allHeadWeights[i][0]; // Token 0 的注意力
console.log(
` Head${i}: [${w.map(v => v.toFixed(3)).join(', ')}] 最大关注 Token${w.indexOf(Math.max(...w))}`
);
}
// Decoder Mask 测试
const causalMask = Array.from({ length: seqLen }, (_, i) =>
Array.from({ length: seqLen }, (_, j) => (j <= i ? 1 : 0))
);
const { allHeadWeights: maskedWeights } = mha.forward(X, causalMask);
console.log('\nDecoder Mask 后 Token3 的注意力:');
console.log(` [${maskedWeights[0][3].map(v => v.toFixed(3)).join(', ')}]`);
console.log(' ↑ Token 4-9 均为 0.000(因果遮蔽成功)');
}
test();
运行结果:
MHA 输入: [10, 512]
MHA 输出: [10, 512]
注意力头数: 8
各头 Token0→所有Token 注意力分布对比:
Head0: [0.098, 0.102, 0.101, 0.099, 0.100, 0.100, 0.099, 0.101, 0.100, 0.099] 最大关注 Token1
Head1: [0.098, 0.098, 0.107, 0.097, 0.100, 0.103, 0.098, 0.098, 0.101, 0.100] 最大关注 Token2
Head2: [0.104, 0.100, 0.099, 0.098, 0.100, 0.098, 0.099, 0.101, 0.102, 0.099] 最大关注 Token0
不同头的注意力模式不同(未训练时为随机差异,训练后会分化为语法/语义等专门模式)。
四、逐行精讲
4.1 为什么是 8 个头,不是 16 或 4?
论文做了消融实验:h=1 → PPL 高(表达能力不足);h=8 和 h=16 效果接近,但 h=8 更快。选择 h 的工程原则:d_model 能被 h 整除(保持 d_k = d_v = d_model/h 为整数),且单头维度不能太小(d_k >= 32 为宜)。
4.2 Concat 之后为什么要再乘 W_O?
// ❌ 不做 W_O:各头输出直接拼接但彼此独立
const concat = [...head0, ...head1, ..., ...head7];
// 问题:头与头之间没有交互!每个头只看到了自己投影的子空间
// ✅ 有 W_O:混合所有头的信息
const output = matMul(concat, this.WO);
// W_O 学习「哪个头的信息对于当前位置更重要」
4.3 Mask 的两种典型用法
Encoder(双向): 所有 token 可关注所有 token → mask = undefined
Decoder(因果/自回归): token i 只能看到 [0, i] → causal mask
// Decoder 的因果遮罩
const causalMask = Array.from({ length: n }, (_, i) =>
Array.from({ length: n }, (_, j) => (j <= i ? 1 : 0))
);
// 效果: 矩阵下三角为 1,上三角为 0
五、常见问题与踩坑记录
Q1:多头注意力的参数量如何计算?
A:每个头有 3 个矩阵(WQ, WK, WV),每个矩阵 d_model × d_k 参数。加上输出投影 W_O(h·d_k × d_model)。总计:h × 3 × d_model × d_k + d_model × h × d_k = 4 × h × d_model × d_k = 4 × d_model²。以 d_model=512 为例,多头的参数 ≈ 4 × 262144 = 1,048,576。
Q2:为什么 h*d_v 必须等于 d_model?
A:为了残差连接。如果 Concat 后的维度不等于 d_model,就无法做 LayerNorm(X + MHA(X))。实际实现中总是保持 h × d_v = d_model。
六、决策框架
什么场景用 Multi-Head Attention?
├── 文本理解(BERT)→ Encoder MHA,无 Mask
│ └── 需要双向上下文 → mask = undefined
│
├── 文本生成(GPT)→ Decoder MHA,Causal Mask
│ └── 只能看左边的 token → causal mask
│
├── 翻译/摘要(T5/BART)→ Encoder + Cross-Attention Decoder
│ └── Cross-Attention: Q=Decoder, K/V=Encoder
│
└── 视觉(ViT)→ Image as tokens, Encoder MHA
└── 将图片切为 16×16 patches,视作 token 序列
七、面试速记
Q:Multi-Head Attention 比单头好在哪?
A:单头只能学到一种注意力模式(如只关注语义)。多头允许模型在不同表示子空间中并行学习不同类型的依赖:语义、句法、位置、指代等。实验证明 h=8 显著优于 h=1(BLEU 提升 2+ 点)。
Q:Multi-Head Attention 的计算量会变成 h 倍吗?
A:不会。虽然头数增到 h,但每个头的维度缩到 d_model/h,总计算量不变:h × O(n² · d_model/h · d_model/h) = O(n² · d_model² / h)。加上 Concat 和 W_O 的 O(n · d_model²),整体比单头无显著额外开销。
八、总结
- 多头注意力不是「并行跑 8 次单头」,而是让 8 个头各自学习不同表示子空间中的注意力模式,最终通过 W_O 混合所有视角。
- Concat + W_O 是关键:没有 W_O,各头信息互不流通;有 W_O,模型学会「什么时候该信任哪个头」。
- Causal Mask 使同一个 MHA 能同时服务 Encoder(无 Mask)和 Decoder(有 Mask),这是 GPT 系列模型的基础构建块。
- 理解 MHA 是理解 MQA(Multi-Query Attention)和 GQA(Grouped-Query Attention)等现代变体的前提——这些变体都是对「如何减少 KV Cache」的工程优化。
Comments 留言讨论
还没有评论,来抢个沙发,聊聊你的看法~